diff --git a/src/blas/backends/cublas/cublas_batch.cpp b/src/blas/backends/cublas/cublas_batch.cpp index 4481195af..943c4944c 100644 --- a/src/blas/backends/cublas/cublas_batch.cpp +++ b/src/blas/backends/cublas/cublas_batch.cpp @@ -25,6 +25,13 @@ namespace oneapi { namespace math { namespace blas { namespace cublas { + +// Row-major dgmm_batch maps to column-major by swapping the side and m/n. +static inline side dgmm_flip_side(side left_right) { + return left_right == oneapi::math::side::left ? oneapi::math::side::right + : oneapi::math::side::left; +} + namespace column_major { // Buffer APIs @@ -110,35 +117,49 @@ void gemv_batch(sycl::queue& queue, transpose transa, int64_t m, int64_t n, throw unimplemented("blas", "gemv_batch", "for column_major layout"); } -void dgmm_batch(sycl::queue& queue, side left_right, int64_t m, int64_t n, - sycl::buffer& a, int64_t lda, int64_t stride_a, sycl::buffer& x, - int64_t incx, int64_t stride_x, sycl::buffer& c, int64_t ldc, - int64_t stride_c, int64_t batch_size) { - throw unimplemented("blas", "dgmm_batch", "for column_major layout"); +// cuBLAS has no batched/strided variant of dgmm (only cublasdgmm), so +// dgmm_batch is implemented as a loop of cublasdgmm calls. See issue #562. +template +inline void dgmm_batch(const char* func_name, Func func, sycl::queue& queue, side left_right, + int64_t m, int64_t n, sycl::buffer& a, int64_t lda, int64_t stride_a, + sycl::buffer& x, int64_t incx, int64_t stride_x, sycl::buffer& c, + int64_t ldc, int64_t stride_c, int64_t batch_size) { + using cuDataType = typename CudaEquivalentType::Type; + queue.submit([&](sycl::handler& cgh) { + auto a_acc = a.template get_access(cgh); + auto x_acc = x.template get_access(cgh); + auto c_acc = c.template get_access(cgh); + onemath_cublas_host_task(cgh, [=](CublasScopedContextHandler& sc) { + auto handle = sc.get_handle(); + auto a_ = sc.get_mem(a_acc); + auto x_ = sc.get_mem(x_acc); + auto c_ = sc.get_mem(c_acc); + cublasStatus_t err; + auto mode = get_cublas_side_mode(left_right); + for (int64_t i = 0; i < batch_size; i++) { + cublas_native_named_func(func_name, func, err, handle, mode, m, n, + a_ + i * stride_a, lda, x_ + i * stride_x, incx, + c_ + i * stride_c, ldc); + } + }); + }); } -void dgmm_batch(sycl::queue& queue, side left_right, int64_t m, int64_t n, - sycl::buffer& a, int64_t lda, int64_t stride_a, - sycl::buffer& x, int64_t incx, int64_t stride_x, - sycl::buffer& c, int64_t ldc, int64_t stride_c, int64_t batch_size) { - throw unimplemented("blas", "dgmm_batch", "for column_major layout"); -} +#define DGMM_STRIDED_BATCH_LAUNCHER(TYPE, CUBLAS_ROUTINE) \ + void dgmm_batch(sycl::queue& queue, side left_right, int64_t m, int64_t n, \ + sycl::buffer& a, int64_t lda, int64_t stride_a, \ + sycl::buffer& x, int64_t incx, int64_t stride_x, \ + sycl::buffer& c, int64_t ldc, int64_t stride_c, int64_t batch_size) { \ + dgmm_batch(#CUBLAS_ROUTINE, CUBLAS_ROUTINE, queue, left_right, m, n, a, lda, stride_a, x, \ + incx, stride_x, c, ldc, stride_c, batch_size); \ + } -void dgmm_batch(sycl::queue& queue, side left_right, int64_t m, int64_t n, - sycl::buffer, 1>& a, int64_t lda, int64_t stride_a, - sycl::buffer, 1>& x, int64_t incx, int64_t stride_x, - sycl::buffer, 1>& c, int64_t ldc, int64_t stride_c, - int64_t batch_size) { - throw unimplemented("blas", "dgmm_batch", "for column_major layout"); -} +DGMM_STRIDED_BATCH_LAUNCHER(float, cublasSdgmm_64) +DGMM_STRIDED_BATCH_LAUNCHER(double, cublasDdgmm_64) +DGMM_STRIDED_BATCH_LAUNCHER(std::complex, cublasCdgmm_64) +DGMM_STRIDED_BATCH_LAUNCHER(std::complex, cublasZdgmm_64) -void dgmm_batch(sycl::queue& queue, side left_right, int64_t m, int64_t n, - sycl::buffer, 1>& a, int64_t lda, int64_t stride_a, - sycl::buffer, 1>& x, int64_t incx, int64_t stride_x, - sycl::buffer, 1>& c, int64_t ldc, int64_t stride_c, - int64_t batch_size) { - throw unimplemented("blas", "dgmm_batch", "for column_major layout"); -} +#undef DGMM_STRIDED_BATCH_LAUNCHER template inline void gemm_batch_impl(sycl::queue& queue, transpose transa, transpose transb, int64_t m, @@ -550,63 +571,104 @@ GEMV_BATCH_LAUNCHER_USM(std::complex, cublasZgemvBatched) #undef GEMV_BATCH_LAUNCHER_USM -sycl::event dgmm_batch(sycl::queue& queue, side left_right, int64_t m, int64_t n, const float* a, - int64_t lda, int64_t stride_a, const float* x, int64_t incx, - int64_t stride_x, float* c, int64_t ldc, int64_t stride_c, - int64_t batch_size, const std::vector& dependencies) { - throw unimplemented("blas", "dgmm_batch", "for column_major layout"); +// USM strided dgmm_batch: loop over cublasdgmm (no native batched variant). +template +inline sycl::event dgmm_batch(const char* func_name, Func func, sycl::queue& queue, side left_right, + int64_t m, int64_t n, const T* a, int64_t lda, int64_t stride_a, + const T* x, int64_t incx, int64_t stride_x, T* c, int64_t ldc, + int64_t stride_c, int64_t batch_size, + const std::vector& dependencies) { + using cuDataType = typename CudaEquivalentType::Type; + auto done = queue.submit([&](sycl::handler& cgh) { + int64_t num_events = dependencies.size(); + for (int64_t i = 0; i < num_events; i++) { + cgh.depends_on(dependencies[i]); + } + onemath_cublas_host_task(cgh, [=](CublasScopedContextHandler& sc) { + auto handle = sc.get_handle(); + auto a_ = reinterpret_cast(a); + auto x_ = reinterpret_cast(x); + auto c_ = reinterpret_cast(c); + cublasStatus_t err; + auto mode = get_cublas_side_mode(left_right); + for (int64_t i = 0; i < batch_size; i++) { + cublas_native_named_func(func_name, func, err, handle, mode, m, n, + a_ + i * stride_a, lda, x_ + i * stride_x, incx, + c_ + i * stride_c, ldc); + } + }); + }); + return done; } -sycl::event dgmm_batch(sycl::queue& queue, side left_right, int64_t m, int64_t n, const double* a, - int64_t lda, int64_t stride_a, const double* x, int64_t incx, - int64_t stride_x, double* c, int64_t ldc, int64_t stride_c, - int64_t batch_size, const std::vector& dependencies) { - throw unimplemented("blas", "dgmm_batch", "for column_major layout"); -} +#define DGMM_STRIDED_BATCH_LAUNCHER_USM(TYPE, CUBLAS_ROUTINE) \ + sycl::event dgmm_batch(sycl::queue& queue, side left_right, int64_t m, int64_t n, \ + const TYPE* a, int64_t lda, int64_t stride_a, const TYPE* x, \ + int64_t incx, int64_t stride_x, TYPE* c, int64_t ldc, int64_t stride_c, \ + int64_t batch_size, const std::vector& dependencies) { \ + return dgmm_batch(#CUBLAS_ROUTINE, CUBLAS_ROUTINE, queue, left_right, m, n, a, lda, \ + stride_a, x, incx, stride_x, c, ldc, stride_c, batch_size, \ + dependencies); \ + } -sycl::event dgmm_batch(sycl::queue& queue, side left_right, int64_t m, int64_t n, - const std::complex* a, int64_t lda, int64_t stride_a, - const std::complex* x, int64_t incx, int64_t stride_x, - std::complex* c, int64_t ldc, int64_t stride_c, int64_t batch_size, - const std::vector& dependencies) { - throw unimplemented("blas", "dgmm_batch", "for column_major layout"); -} +DGMM_STRIDED_BATCH_LAUNCHER_USM(float, cublasSdgmm_64) +DGMM_STRIDED_BATCH_LAUNCHER_USM(double, cublasDdgmm_64) +DGMM_STRIDED_BATCH_LAUNCHER_USM(std::complex, cublasCdgmm_64) +DGMM_STRIDED_BATCH_LAUNCHER_USM(std::complex, cublasZdgmm_64) -sycl::event dgmm_batch(sycl::queue& queue, side left_right, int64_t m, int64_t n, - const std::complex* a, int64_t lda, int64_t stride_a, - const std::complex* x, int64_t incx, int64_t stride_x, - std::complex* c, int64_t ldc, int64_t stride_c, int64_t batch_size, - const std::vector& dependencies) { - throw unimplemented("blas", "dgmm_batch", "for column_major layout"); -} +#undef DGMM_STRIDED_BATCH_LAUNCHER_USM -sycl::event dgmm_batch(sycl::queue& queue, side* left_right, int64_t* m, int64_t* n, - const float** a, int64_t* lda, const float** x, int64_t* incx, float** c, - int64_t* ldc, int64_t group_count, int64_t* groupsize, - const std::vector& dependencies) { - throw unimplemented("blas", "dgmm_batch", "for column_major layout"); +// USM group dgmm_batch: loop over groups and group members calling cublasdgmm. +// flip_side lets the row-major layer reverse each group's side without writing to +// the caller's left_right array, which the spec defines as an input parameter. +template +inline sycl::event dgmm_batch(const char* func_name, Func func, sycl::queue& queue, + side* left_right, int64_t* m, int64_t* n, const T** a, int64_t* lda, + const T** x, int64_t* incx, T** c, int64_t* ldc, int64_t group_count, + int64_t* groupsize, const std::vector& dependencies, + bool flip_side = false) { + using cuDataType = typename CudaEquivalentType::Type; + auto done = queue.submit([&](sycl::handler& cgh) { + int64_t num_events = dependencies.size(); + for (int64_t i = 0; i < num_events; i++) { + cgh.depends_on(dependencies[i]); + } + onemath_cublas_host_task(cgh, [=](CublasScopedContextHandler& sc) { + auto handle = sc.get_handle(); + cublasStatus_t err; + int64_t offset = 0; + for (int64_t i = 0; i < group_count; i++) { + auto mode = + get_cublas_side_mode(flip_side ? dgmm_flip_side(left_right[i]) : left_right[i]); + for (int64_t j = 0; j < groupsize[i]; j++) { + auto a_ = reinterpret_cast(a[offset + j]); + auto x_ = reinterpret_cast(x[offset + j]); + auto c_ = reinterpret_cast(c[offset + j]); + cublas_native_named_func(func_name, func, err, handle, mode, m[i], n[i], a_, + lda[i], x_, incx[i], c_, ldc[i]); + } + offset += groupsize[i]; + } + }); + }); + return done; } -sycl::event dgmm_batch(sycl::queue& queue, side* left_right, int64_t* m, int64_t* n, - const double** a, int64_t* lda, const double** x, int64_t* incx, double** c, - int64_t* ldc, int64_t group_count, int64_t* groupsize, - const std::vector& dependencies) { - throw unimplemented("blas", "dgmm_batch", "for column_major layout"); -} +#define DGMM_GROUP_BATCH_LAUNCHER_USM(TYPE, CUBLAS_ROUTINE) \ + sycl::event dgmm_batch(sycl::queue& queue, side* left_right, int64_t* m, int64_t* n, \ + const TYPE** a, int64_t* lda, const TYPE** x, int64_t* incx, TYPE** c, \ + int64_t* ldc, int64_t group_count, int64_t* groupsize, \ + const std::vector& dependencies) { \ + return dgmm_batch(#CUBLAS_ROUTINE, CUBLAS_ROUTINE, queue, left_right, m, n, a, lda, x, \ + incx, c, ldc, group_count, groupsize, dependencies); \ + } -sycl::event dgmm_batch(sycl::queue& queue, side* left_right, int64_t* m, int64_t* n, - const std::complex** a, int64_t* lda, const std::complex** x, - int64_t* incx, std::complex** c, int64_t* ldc, int64_t group_count, - int64_t* groupsize, const std::vector& dependencies) { - throw unimplemented("blas", "dgmm_batch", "for column_major layout"); -} +DGMM_GROUP_BATCH_LAUNCHER_USM(float, cublasSdgmm_64) +DGMM_GROUP_BATCH_LAUNCHER_USM(double, cublasDdgmm_64) +DGMM_GROUP_BATCH_LAUNCHER_USM(std::complex, cublasCdgmm_64) +DGMM_GROUP_BATCH_LAUNCHER_USM(std::complex, cublasZdgmm_64) -sycl::event dgmm_batch(sycl::queue& queue, side* left_right, int64_t* m, int64_t* n, - const std::complex** a, int64_t* lda, const std::complex** x, - int64_t* incx, std::complex** c, int64_t* ldc, int64_t group_count, - int64_t* groupsize, const std::vector& dependencies) { - throw unimplemented("blas", "dgmm_batch", "for column_major layout"); -} +#undef DGMM_GROUP_BATCH_LAUNCHER_USM template inline sycl::event gemm_batch_strided_usm_impl(sycl::queue& queue, transpose transa, @@ -1162,35 +1224,21 @@ void gemv_batch(sycl::queue& queue, transpose transa, int64_t m, int64_t n, throw unimplemented("blas", "gemv_batch", "for row_major layout"); } -void dgmm_batch(sycl::queue& queue, side left_right, int64_t m, int64_t n, - sycl::buffer& a, int64_t lda, int64_t stride_a, sycl::buffer& x, - int64_t incx, int64_t stride_x, sycl::buffer& c, int64_t ldc, - int64_t stride_c, int64_t batch_size) { - throw unimplemented("blas", "dgmm_batch", "for row_major layout"); -} - -void dgmm_batch(sycl::queue& queue, side left_right, int64_t m, int64_t n, - sycl::buffer& a, int64_t lda, int64_t stride_a, - sycl::buffer& x, int64_t incx, int64_t stride_x, - sycl::buffer& c, int64_t ldc, int64_t stride_c, int64_t batch_size) { - throw unimplemented("blas", "dgmm_batch", "for row_major layout"); -} +#define DGMM_STRIDED_BATCH_LAUNCHER(TYPE) \ + void dgmm_batch(sycl::queue& queue, side left_right, int64_t m, int64_t n, \ + sycl::buffer& a, int64_t lda, int64_t stride_a, \ + sycl::buffer& x, int64_t incx, int64_t stride_x, \ + sycl::buffer& c, int64_t ldc, int64_t stride_c, int64_t batch_size) { \ + column_major::dgmm_batch(queue, dgmm_flip_side(left_right), n, m, a, lda, stride_a, x, \ + incx, stride_x, c, ldc, stride_c, batch_size); \ + } -void dgmm_batch(sycl::queue& queue, side left_right, int64_t m, int64_t n, - sycl::buffer, 1>& a, int64_t lda, int64_t stride_a, - sycl::buffer, 1>& x, int64_t incx, int64_t stride_x, - sycl::buffer, 1>& c, int64_t ldc, int64_t stride_c, - int64_t batch_size) { - throw unimplemented("blas", "dgmm_batch", "for row_major layout"); -} +DGMM_STRIDED_BATCH_LAUNCHER(float) +DGMM_STRIDED_BATCH_LAUNCHER(double) +DGMM_STRIDED_BATCH_LAUNCHER(std::complex) +DGMM_STRIDED_BATCH_LAUNCHER(std::complex) -void dgmm_batch(sycl::queue& queue, side left_right, int64_t m, int64_t n, - sycl::buffer, 1>& a, int64_t lda, int64_t stride_a, - sycl::buffer, 1>& x, int64_t incx, int64_t stride_x, - sycl::buffer, 1>& c, int64_t ldc, int64_t stride_c, - int64_t batch_size) { - throw unimplemented("blas", "dgmm_batch", "for row_major layout"); -} +#undef DGMM_STRIDED_BATCH_LAUNCHER #define GEMM_STRIDED_BATCH_LAUNCHER(TYPE_A, TYPE_B, TYPE_C, TYPE_S) \ void gemm_batch(sycl::queue& queue, transpose transa, transpose transb, int64_t m, int64_t n, \ @@ -1521,63 +1569,39 @@ sycl::event gemv_batch(sycl::queue& queue, transpose* transa, int64_t* m, int64_ throw unimplemented("blas", "gemv_batch", "for row_major layout"); } -sycl::event dgmm_batch(sycl::queue& queue, side left_right, int64_t m, int64_t n, const float* a, - int64_t lda, int64_t stride_a, const float* x, int64_t incx, - int64_t stride_x, float* c, int64_t ldc, int64_t stride_c, - int64_t batch_size, const std::vector& dependencies) { - throw unimplemented("blas", "dgmm_batch", "for row_major layout"); -} - -sycl::event dgmm_batch(sycl::queue& queue, side left_right, int64_t m, int64_t n, const double* a, - int64_t lda, int64_t stride_a, const double* x, int64_t incx, - int64_t stride_x, double* c, int64_t ldc, int64_t stride_c, - int64_t batch_size, const std::vector& dependencies) { - throw unimplemented("blas", "dgmm_batch", "for row_major layout"); -} - -sycl::event dgmm_batch(sycl::queue& queue, side left_right, int64_t m, int64_t n, - const std::complex* a, int64_t lda, int64_t stride_a, - const std::complex* x, int64_t incx, int64_t stride_x, - std::complex* c, int64_t ldc, int64_t stride_c, int64_t batch_size, - const std::vector& dependencies) { - throw unimplemented("blas", "dgmm_batch", "for row_major layout"); -} - -sycl::event dgmm_batch(sycl::queue& queue, side left_right, int64_t m, int64_t n, - const std::complex* a, int64_t lda, int64_t stride_a, - const std::complex* x, int64_t incx, int64_t stride_x, - std::complex* c, int64_t ldc, int64_t stride_c, int64_t batch_size, - const std::vector& dependencies) { - throw unimplemented("blas", "dgmm_batch", "for row_major layout"); -} - -sycl::event dgmm_batch(sycl::queue& queue, side* left_right, int64_t* m, int64_t* n, - const float** a, int64_t* lda, const float** x, int64_t* incx, float** c, - int64_t* ldc, int64_t group_count, int64_t* groupsize, - const std::vector& dependencies) { - throw unimplemented("blas", "dgmm_batch", "for row_major layout"); -} +#define DGMM_STRIDED_BATCH_LAUNCHER_USM(TYPE) \ + sycl::event dgmm_batch(sycl::queue& queue, side left_right, int64_t m, int64_t n, \ + const TYPE* a, int64_t lda, int64_t stride_a, const TYPE* x, \ + int64_t incx, int64_t stride_x, TYPE* c, int64_t ldc, int64_t stride_c, \ + int64_t batch_size, const std::vector& dependencies) { \ + return column_major::dgmm_batch(queue, dgmm_flip_side(left_right), n, m, a, lda, stride_a, \ + x, incx, stride_x, c, ldc, stride_c, batch_size, \ + dependencies); \ + } -sycl::event dgmm_batch(sycl::queue& queue, side* left_right, int64_t* m, int64_t* n, - const double** a, int64_t* lda, const double** x, int64_t* incx, double** c, - int64_t* ldc, int64_t group_count, int64_t* groupsize, - const std::vector& dependencies) { - throw unimplemented("blas", "dgmm_batch", "for row_major layout"); -} +DGMM_STRIDED_BATCH_LAUNCHER_USM(float) +DGMM_STRIDED_BATCH_LAUNCHER_USM(double) +DGMM_STRIDED_BATCH_LAUNCHER_USM(std::complex) +DGMM_STRIDED_BATCH_LAUNCHER_USM(std::complex) + +#undef DGMM_STRIDED_BATCH_LAUNCHER_USM + +#define DGMM_GROUP_BATCH_LAUNCHER_USM(TYPE, CUBLAS_ROUTINE) \ + sycl::event dgmm_batch(sycl::queue& queue, side* left_right, int64_t* m, int64_t* n, \ + const TYPE** a, int64_t* lda, const TYPE** x, int64_t* incx, TYPE** c, \ + int64_t* ldc, int64_t group_count, int64_t* groupsize, \ + const std::vector& dependencies) { \ + return column_major::dgmm_batch(#CUBLAS_ROUTINE, CUBLAS_ROUTINE, queue, left_right, n, m, \ + a, lda, x, incx, c, ldc, group_count, groupsize, \ + dependencies, /*flip_side=*/true); \ + } -sycl::event dgmm_batch(sycl::queue& queue, side* left_right, int64_t* m, int64_t* n, - const std::complex** a, int64_t* lda, const std::complex** x, - int64_t* incx, std::complex** c, int64_t* ldc, int64_t group_count, - int64_t* groupsize, const std::vector& dependencies) { - throw unimplemented("blas", "dgmm_batch", "for row_major layout"); -} +DGMM_GROUP_BATCH_LAUNCHER_USM(float, cublasSdgmm_64) +DGMM_GROUP_BATCH_LAUNCHER_USM(double, cublasDdgmm_64) +DGMM_GROUP_BATCH_LAUNCHER_USM(std::complex, cublasCdgmm_64) +DGMM_GROUP_BATCH_LAUNCHER_USM(std::complex, cublasZdgmm_64) -sycl::event dgmm_batch(sycl::queue& queue, side* left_right, int64_t* m, int64_t* n, - const std::complex** a, int64_t* lda, const std::complex** x, - int64_t* incx, std::complex** c, int64_t* ldc, int64_t group_count, - int64_t* groupsize, const std::vector& dependencies) { - throw unimplemented("blas", "dgmm_batch", "for row_major layout"); -} +#undef DGMM_GROUP_BATCH_LAUNCHER_USM #define GEMM_STRIDED_BATCH_LAUNCHER_USM(TYPE_A, TYPE_B, TYPE_C, TYPE_S) \ sycl::event gemm_batch(sycl::queue& queue, transpose transa, transpose transb, int64_t m, \ diff --git a/src/blas/backends/rocblas/rocblas_batch.cpp b/src/blas/backends/rocblas/rocblas_batch.cpp index b6e550724..0e71d0c3b 100644 --- a/src/blas/backends/rocblas/rocblas_batch.cpp +++ b/src/blas/backends/rocblas/rocblas_batch.cpp @@ -67,6 +67,13 @@ namespace oneapi { namespace math { namespace blas { namespace rocblas { + +// Row-major dgmm_batch maps to column-major by swapping the side and m/n. +static inline side dgmm_flip_side(side left_right) { + return left_right == oneapi::math::side::left ? oneapi::math::side::right + : oneapi::math::side::left; +} + namespace column_major { // Buffer APIs @@ -192,7 +199,6 @@ inline void dgmm_batch(Func func, sycl::queue& queue, side left_right, int64_t m int64_t incx, int64_t stridex, sycl::buffer& c, int64_t ldc, int64_t stridec, int64_t batch_size) { using rocDataType = typename RocEquivalentType::Type; - overflow_check(m, n, lda, ldc, incx, stridea, stridex, stridec, batch_size); queue.submit([&](sycl::handler& cgh) { auto a_acc = a.template get_access(cgh); @@ -211,6 +217,7 @@ inline void dgmm_batch(Func func, sycl::queue& queue, side left_right, int64_t m }); } +// Use the ILP64 (_64) rocBLAS entry points so 64-bit dimensions are supported. #define DGMM_STRIDED_BATCH_LAUNCHER(TYPE, ROCBLAS_ROUTINE) \ void dgmm_batch(sycl::queue& queue, side left_right, int64_t m, int64_t n, \ sycl::buffer& a, int64_t lda, int64_t stridea, \ @@ -220,10 +227,10 @@ inline void dgmm_batch(Func func, sycl::queue& queue, side left_right, int64_t m ldc, stridec, batch_size); \ } -DGMM_STRIDED_BATCH_LAUNCHER(float, rocblas_sdgmm_strided_batched) -DGMM_STRIDED_BATCH_LAUNCHER(double, rocblas_ddgmm_strided_batched) -DGMM_STRIDED_BATCH_LAUNCHER(std::complex, rocblas_cdgmm_strided_batched) -DGMM_STRIDED_BATCH_LAUNCHER(std::complex, rocblas_zdgmm_strided_batched) +DGMM_STRIDED_BATCH_LAUNCHER(float, rocblas_sdgmm_strided_batched_64) +DGMM_STRIDED_BATCH_LAUNCHER(double, rocblas_ddgmm_strided_batched_64) +DGMM_STRIDED_BATCH_LAUNCHER(std::complex, rocblas_cdgmm_strided_batched_64) +DGMM_STRIDED_BATCH_LAUNCHER(std::complex, rocblas_zdgmm_strided_batched_64) #undef DGMM_STRIDED_BATCH_LAUNCHER @@ -763,7 +770,6 @@ inline sycl::event dgmm_batch(Func func, sycl::queue& queue, side left_right, in int64_t stridex, T* c, int64_t ldc, int64_t stridec, int64_t batch_size, const std::vector& dependencies) { using rocDataType = typename RocEquivalentType::Type; - overflow_check(m, n, incx, stridea, stridex, stridec, batch_size); auto done = queue.submit([&](sycl::handler& cgh) { cgh.depends_on(dependencies); @@ -791,22 +797,22 @@ inline sycl::event dgmm_batch(Func func, sycl::queue& queue, side left_right, in stridex, c, ldc, stridec, batch_size, dependencies); \ } -DGMM_STRIDED_BATCH_LAUNCHER_USM(float, rocblas_sdgmm_strided_batched) -DGMM_STRIDED_BATCH_LAUNCHER_USM(double, rocblas_ddgmm_strided_batched) -DGMM_STRIDED_BATCH_LAUNCHER_USM(std::complex, rocblas_cdgmm_strided_batched) -DGMM_STRIDED_BATCH_LAUNCHER_USM(std::complex, rocblas_zdgmm_strided_batched) +DGMM_STRIDED_BATCH_LAUNCHER_USM(float, rocblas_sdgmm_strided_batched_64) +DGMM_STRIDED_BATCH_LAUNCHER_USM(double, rocblas_ddgmm_strided_batched_64) +DGMM_STRIDED_BATCH_LAUNCHER_USM(std::complex, rocblas_cdgmm_strided_batched_64) +DGMM_STRIDED_BATCH_LAUNCHER_USM(std::complex, rocblas_zdgmm_strided_batched_64) #undef DGMM_STRIDED_BATCH_LAUNCHER_USM +// flip_side lets the row-major layer reverse each group's side without writing to +// the caller's left_right array, which the spec defines as an input parameter. template inline sycl::event dgmm_batch(Func func, sycl::queue& queue, side* left_right, int64_t* m, int64_t* n, const T** a, int64_t* lda, const T** x, int64_t* incx, T** c, int64_t* ldc, int64_t group_count, int64_t* group_size, - const std::vector& dependencies) { + const std::vector& dependencies, + bool flip_side = false) { using rocDataType = typename RocEquivalentType::Type; - for (int64_t i = 0; i < group_count; i++) { - overflow_check(m[i], n[i], lda[i], ldc[i], incx[i], group_size[i]); - } auto done = queue.submit([&](sycl::handler& cgh) { cgh.depends_on(dependencies); @@ -819,9 +825,10 @@ inline sycl::event dgmm_batch(Func func, sycl::queue& queue, side* left_right, i auto** a_ = reinterpret_cast(a); auto** x_ = reinterpret_cast(x); auto** c_ = reinterpret_cast(c); - rocblas_native_func(func, err, handle, get_rocblas_side_mode(left_right[i]), - (int)m[i], (int)n[i], a_ + offset, (int)lda[i], x_ + offset, - (int)incx[i], c_ + offset, (int)ldc[i], (int)group_size[i]); + const auto side_i = flip_side ? dgmm_flip_side(left_right[i]) : left_right[i]; + rocblas_native_func(func, err, handle, get_rocblas_side_mode(side_i), m[i], n[i], + a_ + offset, lda[i], x_ + offset, incx[i], c_ + offset, ldc[i], + group_size[i]); offset += group_size[i]; } }); @@ -839,10 +846,10 @@ inline sycl::event dgmm_batch(Func func, sycl::queue& queue, side* left_right, i group_count, group_size, dependencies); \ } -DGMM_BATCH_LAUNCHER_USM(float, rocblas_sdgmm_batched) -DGMM_BATCH_LAUNCHER_USM(double, rocblas_ddgmm_batched) -DGMM_BATCH_LAUNCHER_USM(std::complex, rocblas_cdgmm_batched) -DGMM_BATCH_LAUNCHER_USM(std::complex, rocblas_zdgmm_batched) +DGMM_BATCH_LAUNCHER_USM(float, rocblas_sdgmm_batched_64) +DGMM_BATCH_LAUNCHER_USM(double, rocblas_ddgmm_batched_64) +DGMM_BATCH_LAUNCHER_USM(std::complex, rocblas_cdgmm_batched_64) +DGMM_BATCH_LAUNCHER_USM(std::complex, rocblas_zdgmm_batched_64) #undef DGMM_BATCH_LAUNCHER @@ -1508,11 +1515,8 @@ inline void dgmm_batch(Func func, sycl::queue& queue, side left_right, int64_t m sycl::buffer& a, int64_t lda, int64_t stridea, sycl::buffer& x, int64_t incx, int64_t stridex, sycl::buffer& c, int64_t ldc, int64_t stridec, int64_t batch_size) { - auto new_side = left_right == oneapi::math::side::left ? oneapi::math::side::right - : oneapi::math::side::left; - - column_major::dgmm_batch(func, queue, new_side, n, m, a, lda, stridea, x, incx, stridex, c, ldc, - stridec, batch_size); + column_major::dgmm_batch(func, queue, dgmm_flip_side(left_right), n, m, a, lda, stridea, x, + incx, stridex, c, ldc, stridec, batch_size); } #define DGMM_STRIDED_BATCH_LAUNCHER(TYPE, ROCBLAS_ROUTINE) \ @@ -1524,10 +1528,10 @@ inline void dgmm_batch(Func func, sycl::queue& queue, side left_right, int64_t m ldc, stridec, batch_size); \ } -DGMM_STRIDED_BATCH_LAUNCHER(float, rocblas_sdgmm_strided_batched) -DGMM_STRIDED_BATCH_LAUNCHER(double, rocblas_ddgmm_strided_batched) -DGMM_STRIDED_BATCH_LAUNCHER(std::complex, rocblas_cdgmm_strided_batched) -DGMM_STRIDED_BATCH_LAUNCHER(std::complex, rocblas_zdgmm_strided_batched) +DGMM_STRIDED_BATCH_LAUNCHER(float, rocblas_sdgmm_strided_batched_64) +DGMM_STRIDED_BATCH_LAUNCHER(double, rocblas_ddgmm_strided_batched_64) +DGMM_STRIDED_BATCH_LAUNCHER(std::complex, rocblas_cdgmm_strided_batched_64) +DGMM_STRIDED_BATCH_LAUNCHER(std::complex, rocblas_zdgmm_strided_batched_64) #undef DGMM_STRIDED_BATCH_LAUNCHER @@ -1994,11 +1998,8 @@ inline sycl::event dgmm_batch(Func func, sycl::queue& queue, side left_right, in const T* a, int64_t lda, int64_t stridea, const T* x, int64_t incx, int64_t stridex, T* c, int64_t ldc, int64_t stridec, int64_t batch_size, const std::vector& dependencies) { - auto new_side = left_right == oneapi::math::side::left ? oneapi::math::side::right - : oneapi::math::side::left; - - return column_major::dgmm_batch(func, queue, new_side, n, m, a, lda, stridea, x, incx, stridex, - c, ldc, stridec, batch_size, dependencies); + return column_major::dgmm_batch(func, queue, dgmm_flip_side(left_right), n, m, a, lda, stridea, + x, incx, stridex, c, ldc, stridec, batch_size, dependencies); } #define DGMM_STRIDED_BATCH_LAUNCHER_USM(TYPE, ROCBLAS_ROUTINE) \ @@ -2010,10 +2011,10 @@ inline sycl::event dgmm_batch(Func func, sycl::queue& queue, side left_right, in stridex, c, ldc, stridec, batch_size, dependencies); \ } -DGMM_STRIDED_BATCH_LAUNCHER_USM(float, rocblas_sdgmm_strided_batched) -DGMM_STRIDED_BATCH_LAUNCHER_USM(double, rocblas_ddgmm_strided_batched) -DGMM_STRIDED_BATCH_LAUNCHER_USM(std::complex, rocblas_cdgmm_strided_batched) -DGMM_STRIDED_BATCH_LAUNCHER_USM(std::complex, rocblas_zdgmm_strided_batched) +DGMM_STRIDED_BATCH_LAUNCHER_USM(float, rocblas_sdgmm_strided_batched_64) +DGMM_STRIDED_BATCH_LAUNCHER_USM(double, rocblas_ddgmm_strided_batched_64) +DGMM_STRIDED_BATCH_LAUNCHER_USM(std::complex, rocblas_cdgmm_strided_batched_64) +DGMM_STRIDED_BATCH_LAUNCHER_USM(std::complex, rocblas_zdgmm_strided_batched_64) #undef DGMM_STRIDED_BATCH_LAUNCHER_USM @@ -2022,14 +2023,8 @@ inline sycl::event dgmm_batch(Func func, sycl::queue& queue, side* left_right, i int64_t* n, const T** a, int64_t* lda, const T** x, int64_t* incx, T** c, int64_t* ldc, int64_t group_count, int64_t* group_size, const std::vector& dependencies) { - for (int64_t i = 0; i < group_count; i++) { - const auto new_side = left_right[i] == oneapi::math::side::left ? oneapi::math::side::right - : oneapi::math::side::left; - left_right[i] = new_side; - } - return column_major::dgmm_batch(func, queue, left_right, n, m, a, lda, x, incx, c, ldc, - group_count, group_size, dependencies); + group_count, group_size, dependencies, /*flip_side=*/true); } #define DGMM_BATCH_LAUNCHER_USM(TYPE, ROCBLAS_ROUTINE) \ @@ -2041,10 +2036,10 @@ inline sycl::event dgmm_batch(Func func, sycl::queue& queue, side* left_right, i group_count, group_size, dependencies); \ } -DGMM_BATCH_LAUNCHER_USM(float, rocblas_sdgmm_batched) -DGMM_BATCH_LAUNCHER_USM(double, rocblas_ddgmm_batched) -DGMM_BATCH_LAUNCHER_USM(std::complex, rocblas_cdgmm_batched) -DGMM_BATCH_LAUNCHER_USM(std::complex, rocblas_zdgmm_batched) +DGMM_BATCH_LAUNCHER_USM(float, rocblas_sdgmm_batched_64) +DGMM_BATCH_LAUNCHER_USM(double, rocblas_ddgmm_batched_64) +DGMM_BATCH_LAUNCHER_USM(std::complex, rocblas_cdgmm_batched_64) +DGMM_BATCH_LAUNCHER_USM(std::complex, rocblas_zdgmm_batched_64) #undef DGMM_BATCH_LAUNCHER diff --git a/tests/unit_tests/blas/batch/dgmm_batch_stride.cpp b/tests/unit_tests/blas/batch/dgmm_batch_stride.cpp index 9fa01c52a..d2282a9b7 100644 --- a/tests/unit_tests/blas/batch/dgmm_batch_stride.cpp +++ b/tests/unit_tests/blas/batch/dgmm_batch_stride.cpp @@ -191,6 +191,10 @@ TEST_P(DgmmBatchStrideTests, RealSinglePrecision) { oneapi::math::side::left, -2, 5)); EXPECT_TRUEORSKIP(test(std::get<0>(GetParam()), std::get<1>(GetParam()), oneapi::math::side::left, 1, 5)); + EXPECT_TRUEORSKIP(test(std::get<0>(GetParam()), std::get<1>(GetParam()), + oneapi::math::side::right, 3, 5)); + EXPECT_TRUEORSKIP(test(std::get<0>(GetParam()), std::get<1>(GetParam()), + oneapi::math::side::left, -1, 5)); } TEST_P(DgmmBatchStrideTests, RealDoublePrecision) { @@ -208,6 +212,10 @@ TEST_P(DgmmBatchStrideTests, RealDoublePrecision) { oneapi::math::side::left, -2, 5)); EXPECT_TRUEORSKIP(test(std::get<0>(GetParam()), std::get<1>(GetParam()), oneapi::math::side::left, 1, 5)); + EXPECT_TRUEORSKIP(test(std::get<0>(GetParam()), std::get<1>(GetParam()), + oneapi::math::side::right, 3, 5)); + EXPECT_TRUEORSKIP(test(std::get<0>(GetParam()), std::get<1>(GetParam()), + oneapi::math::side::left, -1, 5)); } TEST_P(DgmmBatchStrideTests, ComplexSinglePrecision) { @@ -223,6 +231,10 @@ TEST_P(DgmmBatchStrideTests, ComplexSinglePrecision) { oneapi::math::side::left, -2, 5)); EXPECT_TRUEORSKIP(test>(std::get<0>(GetParam()), std::get<1>(GetParam()), oneapi::math::side::left, 1, 5)); + EXPECT_TRUEORSKIP(test>(std::get<0>(GetParam()), std::get<1>(GetParam()), + oneapi::math::side::right, 3, 5)); + EXPECT_TRUEORSKIP(test>(std::get<0>(GetParam()), std::get<1>(GetParam()), + oneapi::math::side::left, -1, 5)); } TEST_P(DgmmBatchStrideTests, ComplexDoublePrecision) { @@ -240,6 +252,10 @@ TEST_P(DgmmBatchStrideTests, ComplexDoublePrecision) { oneapi::math::side::left, -2, 5)); EXPECT_TRUEORSKIP(test>(std::get<0>(GetParam()), std::get<1>(GetParam()), oneapi::math::side::left, 1, 5)); + EXPECT_TRUEORSKIP(test>(std::get<0>(GetParam()), std::get<1>(GetParam()), + oneapi::math::side::right, 3, 5)); + EXPECT_TRUEORSKIP(test>(std::get<0>(GetParam()), std::get<1>(GetParam()), + oneapi::math::side::left, -1, 5)); } INSTANTIATE_TEST_SUITE_P(DgmmBatchStrideTestSuite, DgmmBatchStrideTests, diff --git a/tests/unit_tests/blas/batch/dgmm_batch_stride_usm.cpp b/tests/unit_tests/blas/batch/dgmm_batch_stride_usm.cpp index c486ac90e..30c48f8cf 100644 --- a/tests/unit_tests/blas/batch/dgmm_batch_stride_usm.cpp +++ b/tests/unit_tests/blas/batch/dgmm_batch_stride_usm.cpp @@ -196,6 +196,10 @@ TEST_P(DgmmBatchStrideUsmTests, RealSinglePrecision) { oneapi::math::side::left, -2, 5)); EXPECT_TRUEORSKIP(test(std::get<0>(GetParam()), std::get<1>(GetParam()), oneapi::math::side::left, 1, 5)); + EXPECT_TRUEORSKIP(test(std::get<0>(GetParam()), std::get<1>(GetParam()), + oneapi::math::side::right, 3, 5)); + EXPECT_TRUEORSKIP(test(std::get<0>(GetParam()), std::get<1>(GetParam()), + oneapi::math::side::left, -1, 5)); } TEST_P(DgmmBatchStrideUsmTests, RealDoublePrecision) { @@ -213,6 +217,10 @@ TEST_P(DgmmBatchStrideUsmTests, RealDoublePrecision) { oneapi::math::side::left, -2, 5)); EXPECT_TRUEORSKIP(test(std::get<0>(GetParam()), std::get<1>(GetParam()), oneapi::math::side::left, 1, 5)); + EXPECT_TRUEORSKIP(test(std::get<0>(GetParam()), std::get<1>(GetParam()), + oneapi::math::side::right, 3, 5)); + EXPECT_TRUEORSKIP(test(std::get<0>(GetParam()), std::get<1>(GetParam()), + oneapi::math::side::left, -1, 5)); } TEST_P(DgmmBatchStrideUsmTests, ComplexSinglePrecision) { @@ -228,6 +236,10 @@ TEST_P(DgmmBatchStrideUsmTests, ComplexSinglePrecision) { oneapi::math::side::left, -2, 5)); EXPECT_TRUEORSKIP(test>(std::get<0>(GetParam()), std::get<1>(GetParam()), oneapi::math::side::left, 1, 5)); + EXPECT_TRUEORSKIP(test>(std::get<0>(GetParam()), std::get<1>(GetParam()), + oneapi::math::side::right, 3, 5)); + EXPECT_TRUEORSKIP(test>(std::get<0>(GetParam()), std::get<1>(GetParam()), + oneapi::math::side::left, -1, 5)); } TEST_P(DgmmBatchStrideUsmTests, ComplexDoublePrecision) { @@ -245,6 +257,10 @@ TEST_P(DgmmBatchStrideUsmTests, ComplexDoublePrecision) { oneapi::math::side::left, -2, 5)); EXPECT_TRUEORSKIP(test>(std::get<0>(GetParam()), std::get<1>(GetParam()), oneapi::math::side::left, 1, 5)); + EXPECT_TRUEORSKIP(test>(std::get<0>(GetParam()), std::get<1>(GetParam()), + oneapi::math::side::right, 3, 5)); + EXPECT_TRUEORSKIP(test>(std::get<0>(GetParam()), std::get<1>(GetParam()), + oneapi::math::side::left, -1, 5)); } INSTANTIATE_TEST_SUITE_P(DgmmBatchStrideUsmTestSuite, DgmmBatchStrideUsmTests, diff --git a/tests/unit_tests/blas/batch/dgmm_batch_usm.cpp b/tests/unit_tests/blas/batch/dgmm_batch_usm.cpp index 3df3bffd2..230cd80e6 100644 --- a/tests/unit_tests/blas/batch/dgmm_batch_usm.cpp +++ b/tests/unit_tests/blas/batch/dgmm_batch_usm.cpp @@ -185,6 +185,9 @@ int test(device* dev, oneapi::math::layout layout, int64_t group_count) { // Call DPC++ DGMM_BATCH. + // left_right is an input parameter, so the backend must leave it untouched. + std::vector left_right_orig(left_right.begin(), left_right.end()); + try { #ifdef CALL_RT_API switch (layout) { @@ -254,6 +257,11 @@ int test(device* dev, oneapi::math::layout layout, int64_t group_count) { } bool good = true; + if (!std::equal(left_right_orig.begin(), left_right_orig.end(), left_right.begin())) { + std::cout << "Error: DGMM_BATCH overwrote the input left_right array\n"; + good = false; + } + // Compare the results of reference implementation and DPC++ implementation. idx = 0; for (i = 0; i < group_count; i++) { @@ -289,25 +297,37 @@ class DgmmBatchUsmTests : public ::testing::TestWithParam> {}; TEST_P(DgmmBatchUsmTests, RealSinglePrecision) { + EXPECT_TRUEORSKIP(test(std::get<0>(GetParam()), std::get<1>(GetParam()), 1)); EXPECT_TRUEORSKIP(test(std::get<0>(GetParam()), std::get<1>(GetParam()), 5)); + EXPECT_TRUEORSKIP(test(std::get<0>(GetParam()), std::get<1>(GetParam()), 10)); } TEST_P(DgmmBatchUsmTests, RealDoublePrecision) { CHECK_DOUBLE_ON_DEVICE(std::get<0>(GetParam())); + EXPECT_TRUEORSKIP(test(std::get<0>(GetParam()), std::get<1>(GetParam()), 1)); EXPECT_TRUEORSKIP(test(std::get<0>(GetParam()), std::get<1>(GetParam()), 5)); + EXPECT_TRUEORSKIP(test(std::get<0>(GetParam()), std::get<1>(GetParam()), 10)); } TEST_P(DgmmBatchUsmTests, ComplexSinglePrecision) { + EXPECT_TRUEORSKIP( + test>(std::get<0>(GetParam()), std::get<1>(GetParam()), 1)); EXPECT_TRUEORSKIP( test>(std::get<0>(GetParam()), std::get<1>(GetParam()), 5)); + EXPECT_TRUEORSKIP( + test>(std::get<0>(GetParam()), std::get<1>(GetParam()), 10)); } TEST_P(DgmmBatchUsmTests, ComplexDoublePrecision) { CHECK_DOUBLE_ON_DEVICE(std::get<0>(GetParam())); + EXPECT_TRUEORSKIP( + test>(std::get<0>(GetParam()), std::get<1>(GetParam()), 1)); EXPECT_TRUEORSKIP( test>(std::get<0>(GetParam()), std::get<1>(GetParam()), 5)); + EXPECT_TRUEORSKIP( + test>(std::get<0>(GetParam()), std::get<1>(GetParam()), 10)); } INSTANTIATE_TEST_SUITE_P(DgmmBatchUsmTestSuite, DgmmBatchUsmTests,