From d46d55bc581694834429e2dc39d91295dca186d8 Mon Sep 17 00:00:00 2001 From: Zheming Jin Date: Tue, 11 Aug 2026 14:40:58 +0000 Subject: [PATCH 1/4] [cublas] Implement dgmm_batch backend cuBLAS exposes only the non-batched cublasdgmm (no strided/batched variant exists in any CUDA library), so the cuBLAS backend previously threw unimplemented for all dgmm_batch entry points (issue #599 -> #562). Implement dgmm_batch for the cuBLAS backend by looping over cublasdgmm: - Buffer strided, USM strided, and USM group APIs - float, double, complex, complex - Row-major handled by flipping the side and swapping m/n, then delegating to the column-major implementation (matches the rocBLAS backend) Also extend the shared dgmm_batch tests with additional incx values (3 and -1) and group_count values (1 and 10) for more rigorous coverage. Verified on an NVIDIA RTX PRO 6000 Blackwell (sm_120), CUDA 13.3 / cuBLAS 13.6: all DgmmBatch{,Stride}{,Usm} tests pass (col/row major, all four types, run-time and compile-time dispatch). Co-authored-by: Cursor --- src/blas/backends/cublas/cublas_batch.cpp | 328 ++++++++++-------- .../blas/batch/dgmm_batch_stride.cpp | 16 + .../blas/batch/dgmm_batch_stride_usm.cpp | 16 + .../unit_tests/blas/batch/dgmm_batch_usm.cpp | 12 + 4 files changed, 220 insertions(+), 152 deletions(-) diff --git a/src/blas/backends/cublas/cublas_batch.cpp b/src/blas/backends/cublas/cublas_batch.cpp index 4481195af..3d32d7a59 100644 --- a/src/blas/backends/cublas/cublas_batch.cpp +++ b/src/blas/backends/cublas/cublas_batch.cpp @@ -110,35 +110,50 @@ 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; + overflow_check(m, n, lda, ldc, stride_a, stride_x, stride_c, batch_size); + 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, (int)m, (int)n, + a_ + i * stride_a, (int)lda, x_ + i * stride_x, (int)incx, + c_ + i * stride_c, (int)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) +DGMM_STRIDED_BATCH_LAUNCHER(double, cublasDdgmm) +DGMM_STRIDED_BATCH_LAUNCHER(std::complex, cublasCdgmm) +DGMM_STRIDED_BATCH_LAUNCHER(std::complex, cublasZdgmm) -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 +565,103 @@ 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; + overflow_check(m, n, lda, ldc, stride_a, stride_x, stride_c, batch_size); + 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, (int)m, (int)n, + a_ + i * stride_a, (int)lda, x_ + i * stride_x, (int)incx, + c_ + i * stride_c, (int)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) +DGMM_STRIDED_BATCH_LAUNCHER_USM(double, cublasDdgmm) +DGMM_STRIDED_BATCH_LAUNCHER_USM(std::complex, cublasCdgmm) +DGMM_STRIDED_BATCH_LAUNCHER_USM(std::complex, cublasZdgmm) -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. +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) { + using cuDataType = typename CudaEquivalentType::Type; + for (int64_t i = 0; i < group_count; i++) { + overflow_check(m[i], n[i], lda[i], ldc[i], groupsize[i]); + } + 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(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, (int)m[i], (int)n[i], + a_, (int)lda[i], x_, (int)incx[i], c_, (int)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) +DGMM_GROUP_BATCH_LAUNCHER_USM(double, cublasDdgmm) +DGMM_GROUP_BATCH_LAUNCHER_USM(std::complex, cublasCdgmm) +DGMM_GROUP_BATCH_LAUNCHER_USM(std::complex, cublasZdgmm) -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 +1217,27 @@ 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"); +// 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; } -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 +1568,40 @@ 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) \ + 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) { \ + for (int64_t i = 0; i < group_count; i++) \ + left_right[i] = dgmm_flip_side(left_right[i]); \ + return column_major::dgmm_batch(queue, left_right, n, m, 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 row_major layout"); -} +DGMM_GROUP_BATCH_LAUNCHER_USM(float) +DGMM_GROUP_BATCH_LAUNCHER_USM(double) +DGMM_GROUP_BATCH_LAUNCHER_USM(std::complex) +DGMM_GROUP_BATCH_LAUNCHER_USM(std::complex) -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/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..71ca059b4 100644 --- a/tests/unit_tests/blas/batch/dgmm_batch_usm.cpp +++ b/tests/unit_tests/blas/batch/dgmm_batch_usm.cpp @@ -289,25 +289,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, From bb5fb69eec0e630778000d498139046af526334f Mon Sep 17 00:00:00 2001 From: Zheming Jin Date: Tue, 11 Aug 2026 14:53:54 +0000 Subject: [PATCH 2/4] [cublas] Use cublasdgmm_64 for full 64-bit dimension support Switch dgmm_batch to the cublasdgmm_64 entry points and pass int64_t m/n/lda/incx/ldc directly (no int casts), and drop the 32-bit overflow_check so dimensions beyond 2^31 are supported. Co-authored-by: Cursor --- src/blas/backends/cublas/cublas_batch.cpp | 45 ++++++++++------------- 1 file changed, 20 insertions(+), 25 deletions(-) diff --git a/src/blas/backends/cublas/cublas_batch.cpp b/src/blas/backends/cublas/cublas_batch.cpp index 3d32d7a59..ca4168c3f 100644 --- a/src/blas/backends/cublas/cublas_batch.cpp +++ b/src/blas/backends/cublas/cublas_batch.cpp @@ -118,7 +118,6 @@ inline void dgmm_batch(const char* func_name, Func func, sycl::queue& queue, sid 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; - overflow_check(m, n, lda, ldc, stride_a, stride_x, stride_c, batch_size); queue.submit([&](sycl::handler& cgh) { auto a_acc = a.template get_access(cgh); auto x_acc = x.template get_access(cgh); @@ -131,9 +130,9 @@ inline void dgmm_batch(const char* func_name, Func func, sycl::queue& queue, sid 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, (int)m, (int)n, - a_ + i * stride_a, (int)lda, x_ + i * stride_x, (int)incx, - c_ + i * stride_c, (int)ldc); + 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); } }); }); @@ -148,10 +147,10 @@ inline void dgmm_batch(const char* func_name, Func func, sycl::queue& queue, sid incx, stride_x, c, ldc, stride_c, batch_size); \ } -DGMM_STRIDED_BATCH_LAUNCHER(float, cublasSdgmm) -DGMM_STRIDED_BATCH_LAUNCHER(double, cublasDdgmm) -DGMM_STRIDED_BATCH_LAUNCHER(std::complex, cublasCdgmm) -DGMM_STRIDED_BATCH_LAUNCHER(std::complex, cublasZdgmm) +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) #undef DGMM_STRIDED_BATCH_LAUNCHER @@ -573,7 +572,6 @@ inline sycl::event dgmm_batch(const char* func_name, Func func, sycl::queue& que int64_t stride_c, int64_t batch_size, const std::vector& dependencies) { using cuDataType = typename CudaEquivalentType::Type; - overflow_check(m, n, lda, ldc, stride_a, stride_x, stride_c, batch_size); auto done = queue.submit([&](sycl::handler& cgh) { int64_t num_events = dependencies.size(); for (int64_t i = 0; i < num_events; i++) { @@ -587,9 +585,9 @@ inline sycl::event dgmm_batch(const char* func_name, Func func, sycl::queue& que 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, (int)m, (int)n, - a_ + i * stride_a, (int)lda, x_ + i * stride_x, (int)incx, - c_ + i * stride_c, (int)ldc); + 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); } }); }); @@ -605,10 +603,10 @@ inline sycl::event dgmm_batch(const char* func_name, Func func, sycl::queue& que stride_a, x, incx, stride_x, c, ldc, stride_c, batch_size, dependencies); \ } -DGMM_STRIDED_BATCH_LAUNCHER_USM(float, cublasSdgmm) -DGMM_STRIDED_BATCH_LAUNCHER_USM(double, cublasDdgmm) -DGMM_STRIDED_BATCH_LAUNCHER_USM(std::complex, cublasCdgmm) -DGMM_STRIDED_BATCH_LAUNCHER_USM(std::complex, cublasZdgmm) +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) #undef DGMM_STRIDED_BATCH_LAUNCHER_USM @@ -619,9 +617,6 @@ inline sycl::event dgmm_batch(const char* func_name, Func func, sycl::queue& que int64_t* incx, T** c, int64_t* ldc, int64_t group_count, int64_t* groupsize, const std::vector& dependencies) { using cuDataType = typename CudaEquivalentType::Type; - for (int64_t i = 0; i < group_count; i++) { - overflow_check(m[i], n[i], lda[i], ldc[i], groupsize[i]); - } auto done = queue.submit([&](sycl::handler& cgh) { int64_t num_events = dependencies.size(); for (int64_t i = 0; i < num_events; i++) { @@ -637,8 +632,8 @@ inline sycl::event dgmm_batch(const char* func_name, Func func, sycl::queue& que 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, (int)m[i], (int)n[i], - a_, (int)lda[i], x_, (int)incx[i], c_, (int)ldc[i]); + 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]; } @@ -656,10 +651,10 @@ inline sycl::event dgmm_batch(const char* func_name, Func func, sycl::queue& que incx, c, ldc, group_count, groupsize, dependencies); \ } -DGMM_GROUP_BATCH_LAUNCHER_USM(float, cublasSdgmm) -DGMM_GROUP_BATCH_LAUNCHER_USM(double, cublasDdgmm) -DGMM_GROUP_BATCH_LAUNCHER_USM(std::complex, cublasCdgmm) -DGMM_GROUP_BATCH_LAUNCHER_USM(std::complex, cublasZdgmm) +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) #undef DGMM_GROUP_BATCH_LAUNCHER_USM From 509b748eb41a74b939328356cc7fcd855cd8ae5f Mon Sep 17 00:00:00 2001 From: Zheming Jin Date: Tue, 11 Aug 2026 08:07:34 -0700 Subject: [PATCH 3/4] [rocblas] Use rocblas_?dgmm*_64 for full 64-bit dimension support Switch dgmm_batch (strided buffer, strided USM, and grouped USM) to the ILP64 (_64) rocBLAS entry points and drop the overflow_check calls that capped dimensions below 2^31. Dimensions are now passed straight through as int64_t, matching the cuBLAS backend change. Verified on AMD Instinct MI300A (gfx942): all 24 dgmm_batch CT and RT tests pass. Co-authored-by: Cursor --- src/blas/backends/rocblas/rocblas_batch.cpp | 60 ++++++++++----------- 1 file changed, 28 insertions(+), 32 deletions(-) diff --git a/src/blas/backends/rocblas/rocblas_batch.cpp b/src/blas/backends/rocblas/rocblas_batch.cpp index b6e550724..fd2d8418b 100644 --- a/src/blas/backends/rocblas/rocblas_batch.cpp +++ b/src/blas/backends/rocblas/rocblas_batch.cpp @@ -192,7 +192,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 +210,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 +220,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 +763,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,10 +790,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 @@ -804,9 +803,6 @@ inline sycl::event dgmm_batch(Func func, sycl::queue& queue, side* left_right, i T** c, int64_t* ldc, int64_t group_count, int64_t* group_size, const std::vector& dependencies) { 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 +815,9 @@ 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]); + rocblas_native_func(func, err, handle, get_rocblas_side_mode(left_right[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 +835,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 @@ -1524,10 +1520,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 @@ -2010,10 +2006,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 @@ -2041,10 +2037,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 From eed751d07b98cb65cc958b0c28e6f4528713ee7b Mon Sep 17 00:00:00 2001 From: Zheming Jin Date: Wed, 12 Aug 2026 11:31:56 -0700 Subject: [PATCH 4/4] [cublas][rocblas] Keep dgmm_batch left_right input array unmodified The row-major group dgmm_batch flipped each group's side in place, writing to the caller's left_right array even though the spec defines it as an input parameter. Flip the side when reading it inside the column-major kernel loop instead, and assert in the tests that the array survives the call unchanged. Also reformats the touched dgmm_batch macros so the clang-format check passes. Co-authored-by: Cursor --- src/blas/backends/cublas/cublas_batch.cpp | 79 ++++++++++--------- src/blas/backends/rocblas/rocblas_batch.cpp | 41 +++++----- .../unit_tests/blas/batch/dgmm_batch_usm.cpp | 8 ++ 3 files changed, 70 insertions(+), 58 deletions(-) diff --git a/src/blas/backends/cublas/cublas_batch.cpp b/src/blas/backends/cublas/cublas_batch.cpp index ca4168c3f..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 @@ -594,13 +601,14 @@ inline sycl::event dgmm_batch(const char* func_name, Func func, sycl::queue& que return done; } -#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); \ +#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); \ } DGMM_STRIDED_BATCH_LAUNCHER_USM(float, cublasSdgmm_64) @@ -611,11 +619,14 @@ DGMM_STRIDED_BATCH_LAUNCHER_USM(std::complex, cublasZdgmm_64) #undef DGMM_STRIDED_BATCH_LAUNCHER_USM // 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) { +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(); @@ -627,7 +638,8 @@ inline sycl::event dgmm_batch(const char* func_name, Func func, sycl::queue& que cublasStatus_t err; int64_t offset = 0; for (int64_t i = 0; i < group_count; i++) { - auto mode = get_cublas_side_mode(left_right[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]); @@ -1212,19 +1224,13 @@ void gemv_batch(sycl::queue& queue, transpose transa, int64_t m, int64_t n, throw unimplemented("blas", "gemv_batch", "for row_major layout"); } -// 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; -} - #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); \ + 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); \ } DGMM_STRIDED_BATCH_LAUNCHER(float) @@ -1563,14 +1569,14 @@ sycl::event gemv_batch(sycl::queue& queue, transpose* transa, int64_t* m, int64_ throw unimplemented("blas", "gemv_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); \ +#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); \ } DGMM_STRIDED_BATCH_LAUNCHER_USM(float) @@ -1580,21 +1586,20 @@ DGMM_STRIDED_BATCH_LAUNCHER_USM(std::complex) #undef DGMM_STRIDED_BATCH_LAUNCHER_USM -#define DGMM_GROUP_BATCH_LAUNCHER_USM(TYPE) \ +#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) { \ - for (int64_t i = 0; i < group_count; i++) \ - left_right[i] = dgmm_flip_side(left_right[i]); \ - return column_major::dgmm_batch(queue, left_right, n, m, a, lda, x, incx, c, ldc, \ - group_count, groupsize, 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); \ } -DGMM_GROUP_BATCH_LAUNCHER_USM(float) -DGMM_GROUP_BATCH_LAUNCHER_USM(double) -DGMM_GROUP_BATCH_LAUNCHER_USM(std::complex) -DGMM_GROUP_BATCH_LAUNCHER_USM(std::complex) +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) #undef DGMM_GROUP_BATCH_LAUNCHER_USM diff --git a/src/blas/backends/rocblas/rocblas_batch.cpp b/src/blas/backends/rocblas/rocblas_batch.cpp index fd2d8418b..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 @@ -797,11 +804,14 @@ DGMM_STRIDED_BATCH_LAUNCHER_USM(std::complex, rocblas_zdgmm_strided_batc #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; auto done = queue.submit([&](sycl::handler& cgh) { @@ -815,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]), m[i], - n[i], a_ + offset, lda[i], x_ + offset, incx[i], c_ + offset, - ldc[i], 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]; } }); @@ -1504,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) \ @@ -1990,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) \ @@ -2018,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) \ diff --git a/tests/unit_tests/blas/batch/dgmm_batch_usm.cpp b/tests/unit_tests/blas/batch/dgmm_batch_usm.cpp index 71ca059b4..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++) {