diff --git a/Grid/algorithms/blas/BatchedBlas.h b/Grid/algorithms/blas/BatchedBlas.h index 9616e5795..a900ee429 100644 --- a/Grid/algorithms/blas/BatchedBlas.h +++ b/Grid/algorithms/blas/BatchedBlas.h @@ -695,6 +695,228 @@ public: RealD bytes = 1.0*sizeof(ComplexF)*(m*k+k*n+m*n)*batchCount; } + /////////////////////////////////////////////////////////////////////////////////// + // Explicit-leading-dimension complex single GEMM. + // + // A,B,C may be SLICES of larger parent allocations: lda/ldb/ldc are the + // PARENT strides (>= the compact values the ten-argument overload derives). + // Motivating use: software split-K for tiny-output/huge-K dense multiplies + // (arXiv:2409.03904 fig 11; cf MultiRHSBlockCGLinalg) -- batch over K-chunks + // of a dense slab by pointer offset j*Kchunk with lda = the full K extent, + // then reduce the partial C's. Backends pass lda straight through; only + // the compact overload invented them. + /////////////////////////////////////////////////////////////////////////////////// + void gemmBatched(GridBLASOperation_t OpA, + GridBLASOperation_t OpB, + int m,int n, int k, + ComplexF alpha, + deviceVector &Amk, int lda, + deviceVector &Bkn, int ldb, + ComplexF beta, + deviceVector &Cmn, int ldc, + GridBLASPrecision_t precision = GridBLAS_PRECISION_DEFAULT) + { + RealD t2=usecond(); + int32_t batchCount = Amk.size(); + + GRID_ASSERT( lda >= ((OpA==GridBLAS_OP_N) ? m : k) ); + GRID_ASSERT( ldb >= ((OpB==GridBLAS_OP_N) ? k : n) ); + GRID_ASSERT( ldc >= m ); + + static deviceVector alpha_p(1); + static deviceVector beta_p(1); + acceleratorCopyToDevice((void *)&alpha,(void *)&alpha_p[0],sizeof(ComplexF)); + acceleratorCopyToDevice((void *)&beta ,(void *)&beta_p[0],sizeof(ComplexF)); + RealD t0=usecond(); + + GRID_ASSERT(Bkn.size()==batchCount); + GRID_ASSERT(Cmn.size()==batchCount); +#ifdef GRID_HIP + GRID_ASSERT(precision == GridBLAS_PRECISION_DEFAULT); + hipblasOperation_t hOpA; + hipblasOperation_t hOpB; + if ( OpA == GridBLAS_OP_N ) hOpA = HIPBLAS_OP_N; + if ( OpA == GridBLAS_OP_T ) hOpA = HIPBLAS_OP_T; + if ( OpA == GridBLAS_OP_C ) hOpA = HIPBLAS_OP_C; + if ( OpB == GridBLAS_OP_N ) hOpB = HIPBLAS_OP_N; + if ( OpB == GridBLAS_OP_T ) hOpB = HIPBLAS_OP_T; + if ( OpB == GridBLAS_OP_C ) hOpB = HIPBLAS_OP_C; +#if defined(HIP_VERSION_MAJOR) && (HIP_VERSION_MAJOR >=7) + auto err = hipblasCgemmBatched(gridblasHandle, + hOpA, + hOpB, + m,n,k, + (hipComplex *) &alpha_p[0], + (hipComplex **)&Amk[0], lda, + (hipComplex **)&Bkn[0], ldb, + (hipComplex *) &beta_p[0], + (hipComplex **)&Cmn[0], ldc, + batchCount); +#else + auto err = hipblasCgemmBatched(gridblasHandle, + hOpA, + hOpB, + m,n,k, + (hipblasComplex *) &alpha_p[0], + (hipblasComplex **)&Amk[0], lda, + (hipblasComplex **)&Bkn[0], ldb, + (hipblasComplex *) &beta_p[0], + (hipblasComplex **)&Cmn[0], ldc, + batchCount); +#endif + GRID_ASSERT(err==HIPBLAS_STATUS_SUCCESS); +#endif +#ifdef GRID_CUDA + cublasOperation_t hOpA; + cublasOperation_t hOpB; + if ( OpA == GridBLAS_OP_N ) hOpA = CUBLAS_OP_N; + if ( OpA == GridBLAS_OP_T ) hOpA = CUBLAS_OP_T; + if ( OpA == GridBLAS_OP_C ) hOpA = CUBLAS_OP_C; + if ( OpB == GridBLAS_OP_N ) hOpB = CUBLAS_OP_N; + if ( OpB == GridBLAS_OP_T ) hOpB = CUBLAS_OP_T; + if ( OpB == GridBLAS_OP_C ) hOpB = CUBLAS_OP_C; + cublasStatus_t err; + if (precision == GridBLAS_PRECISION_DEFAULT) { + err = cublasCgemmBatched(gridblasHandle, + hOpA, + hOpB, + m,n,k, + (cuComplex *) &alpha_p[0], + (cuComplex **)&Amk[0], lda, + (cuComplex **)&Bkn[0], ldb, + (cuComplex *) &beta_p[0], + (cuComplex **)&Cmn[0], ldc, + batchCount); + } else { + cublasComputeType_t compute_precision = toDataType(precision); + err = cublasGemmBatchedEx(gridblasHandle, + hOpA, + hOpB, + m,n,k, + (void *) &alpha_p[0], + (void **)&Amk[0], CUDA_C_32F, lda, + (void **)&Bkn[0], CUDA_C_32F, ldb, + (void *) &beta_p[0], + (void **)&Cmn[0], CUDA_C_32F, ldc, + batchCount, compute_precision, CUBLAS_GEMM_DEFAULT); + } + GRID_ASSERT(err==CUBLAS_STATUS_SUCCESS); +#endif +#ifdef GRID_SYCL + GRID_ASSERT(precision == GridBLAS_PRECISION_DEFAULT); + int64_t m64=m; + int64_t n64=n; + int64_t k64=k; + int64_t lda64=lda; + int64_t ldb64=ldb; + int64_t ldc64=ldc; + int64_t batchCount64=batchCount; + + oneapi::mkl::transpose iOpA; + oneapi::mkl::transpose iOpB; + + if ( OpA == GridBLAS_OP_N ) iOpA = oneapi::mkl::transpose::N; + if ( OpA == GridBLAS_OP_T ) iOpA = oneapi::mkl::transpose::T; + if ( OpA == GridBLAS_OP_C ) iOpA = oneapi::mkl::transpose::C; + if ( OpB == GridBLAS_OP_N ) iOpB = oneapi::mkl::transpose::N; + if ( OpB == GridBLAS_OP_T ) iOpB = oneapi::mkl::transpose::T; + if ( OpB == GridBLAS_OP_C ) iOpB = oneapi::mkl::transpose::C; + + oneapi::mkl::blas::column_major::gemm_batch(*gridblasHandle, + &iOpA, + &iOpB, + &m64,&n64,&k64, + (ComplexF *) &alpha_p[0], + (const ComplexF **)&Amk[0], (const int64_t *)&lda64, + (const ComplexF **)&Bkn[0], (const int64_t *)&ldb64, + (ComplexF *) &beta_p[0], + (ComplexF **)&Cmn[0], (const int64_t *)&ldc64, + (int64_t)1,&batchCount64,std::vector()); + synchronise(); +#endif +#if !defined(GRID_SYCL) && !defined(GRID_CUDA) && !defined(GRID_HIP) + GRID_ASSERT(precision == GridBLAS_PRECISION_DEFAULT); + // Reference implementation: Eigen with explicit outer stride + typedef Eigen::Map > eMat; + if ( (OpA == GridBLAS_OP_N ) && (OpB == GridBLAS_OP_N) ) { + thread_for (p, batchCount, { + eMat eAmk(Amk[p],m,k,Eigen::OuterStride<>(lda)); + eMat eBkn(Bkn[p],k,n,Eigen::OuterStride<>(ldb)); + eMat eCmn(Cmn[p],m,n,Eigen::OuterStride<>(ldc)); + if (std::abs(beta) != 0.0) + eCmn = beta * eCmn + alpha * eAmk * eBkn ; + else + eCmn = alpha * eAmk * eBkn ; + }); + } else if ( (OpA == GridBLAS_OP_C ) && (OpB == GridBLAS_OP_N) ) { + thread_for (p, batchCount, { + eMat eAmk(Amk[p],k,m,Eigen::OuterStride<>(lda)); + eMat eBkn(Bkn[p],k,n,Eigen::OuterStride<>(ldb)); + eMat eCmn(Cmn[p],m,n,Eigen::OuterStride<>(ldc)); + if (std::abs(beta) != 0.0) + eCmn = beta * eCmn + alpha * eAmk.adjoint() * eBkn ; + else + eCmn = alpha * eAmk.adjoint() * eBkn ; + }); + } else if ( (OpA == GridBLAS_OP_T ) && (OpB == GridBLAS_OP_N) ) { + thread_for (p, batchCount, { + eMat eAmk(Amk[p],k,m,Eigen::OuterStride<>(lda)); + eMat eBkn(Bkn[p],k,n,Eigen::OuterStride<>(ldb)); + eMat eCmn(Cmn[p],m,n,Eigen::OuterStride<>(ldc)); + if (std::abs(beta) != 0.0) + eCmn = beta * eCmn + alpha * eAmk.transpose() * eBkn ; + else + eCmn = alpha * eAmk.transpose() * eBkn ; + }); + } else if ( (OpA == GridBLAS_OP_N ) && (OpB == GridBLAS_OP_C) ) { + thread_for (p, batchCount, { + eMat eAmk(Amk[p],m,k,Eigen::OuterStride<>(lda)); + eMat eBkn(Bkn[p],n,k,Eigen::OuterStride<>(ldb)); + eMat eCmn(Cmn[p],m,n,Eigen::OuterStride<>(ldc)); + if (std::abs(beta) != 0.0) + eCmn = beta * eCmn + alpha * eAmk * eBkn.adjoint() ; + else + eCmn = alpha * eAmk * eBkn.adjoint() ; + }); + } else if ( (OpA == GridBLAS_OP_N ) && (OpB == GridBLAS_OP_T) ) { + thread_for (p, batchCount, { + eMat eAmk(Amk[p],m,k,Eigen::OuterStride<>(lda)); + eMat eBkn(Bkn[p],n,k,Eigen::OuterStride<>(ldb)); + eMat eCmn(Cmn[p],m,n,Eigen::OuterStride<>(ldc)); + if (std::abs(beta) != 0.0) + eCmn = beta * eCmn + alpha * eAmk * eBkn.transpose() ; + else + eCmn = alpha * eAmk * eBkn.transpose() ; + }); + } else if ( (OpA == GridBLAS_OP_C ) && (OpB == GridBLAS_OP_C) ) { + thread_for (p, batchCount, { + eMat eAmk(Amk[p],k,m,Eigen::OuterStride<>(lda)); + eMat eBkn(Bkn[p],n,k,Eigen::OuterStride<>(ldb)); + eMat eCmn(Cmn[p],m,n,Eigen::OuterStride<>(ldc)); + if (std::abs(beta) != 0.0) + eCmn = beta * eCmn + alpha * eAmk.adjoint() * eBkn.adjoint() ; + else + eCmn = alpha * eAmk.adjoint() * eBkn.adjoint() ; + } ); + } else if ( (OpA == GridBLAS_OP_T ) && (OpB == GridBLAS_OP_T) ) { + thread_for (p, batchCount, { + eMat eAmk(Amk[p],k,m,Eigen::OuterStride<>(lda)); + eMat eBkn(Bkn[p],n,k,Eigen::OuterStride<>(ldb)); + eMat eCmn(Cmn[p],m,n,Eigen::OuterStride<>(ldc)); + if (std::abs(beta) != 0.0) + eCmn = beta * eCmn + alpha * eAmk.transpose() * eBkn.transpose() ; + else + eCmn = alpha * eAmk.transpose() * eBkn.transpose() ; + } ); + } else { + assert(0); + } +#endif + RealD t1=usecond(); + RealD flops = 8.0*m*n*k*batchCount; + RealD bytes = 1.0*sizeof(ComplexF)*(m*k+k*n+m*n)*batchCount; + } + /////////////////////////////////////////////////////////////////////////// // Single precision real GEMM ///////////////////////////////////////////////////////////////////////////