From a7160ac51352d49d31909899320747e1342324f8 Mon Sep 17 00:00:00 2001 From: Peter Boyle Date: Fri, 14 Aug 2026 17:47:59 -0400 Subject: [PATCH] Support for distributed Schur inverse --- Grid/algorithms/blas/BatchedBlas.h | 227 +++++++++++++++++++++++++++++ 1 file changed, 227 insertions(+) diff --git a/Grid/algorithms/blas/BatchedBlas.h b/Grid/algorithms/blas/BatchedBlas.h index d7c4ce98c..0dbd17533 100644 --- a/Grid/algorithms/blas/BatchedBlas.h +++ b/Grid/algorithms/blas/BatchedBlas.h @@ -991,6 +991,233 @@ public: RealD bytes = 1.0*sizeof(ComplexF)*(m*k+k*n+m*n)*batchCount; } + /////////////////////////////////////////////////////////////////////////////////// + // Explicit-leading-dimension complex double GEMM. Mirror of the ComplexF + // overload above; motivating use is the fp64 distributed recursive Schur + // inversion (RecursiveSchurInverse), whose operands are column windows of + // larger row-slab allocations. + /////////////////////////////////////////////////////////////////////////////////// + void gemmBatched(GridBLASOperation_t OpA, + GridBLASOperation_t OpB, + int m,int n, int k, + ComplexD alpha, + deviceVector &Amk, int lda, + deviceVector &Bkn, int ldb, + ComplexD beta, + deviceVector &Cmn, int ldc) + { + 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 ); + + // Cached device constants: copy only on value change (see GridBLASDeviceConstant) + static GridBLASDeviceConstant alpha_c; + static GridBLASDeviceConstant beta_c; + ComplexD *alpha_p = alpha_c.put(alpha); + ComplexD *beta_p = beta_c.put(beta); + RealD t0=usecond(); + + GRID_ASSERT(Bkn.size()==batchCount); + GRID_ASSERT(Cmn.size()==batchCount); +#ifdef GRID_HIP + 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 = hipblasZgemmBatched(gridblasHandle, + hOpA, + hOpB, + m,n,k, + (hipDoubleComplex *) &alpha_p[0], + (hipDoubleComplex **)&Amk[0], lda, + (hipDoubleComplex **)&Bkn[0], ldb, + (hipDoubleComplex *) &beta_p[0], + (hipDoubleComplex **)&Cmn[0], ldc, + batchCount); +#else + auto err = hipblasZgemmBatched(gridblasHandle, + hOpA, + hOpB, + m,n,k, + (hipblasDoubleComplex *) &alpha_p[0], + (hipblasDoubleComplex **)&Amk[0], lda, + (hipblasDoubleComplex **)&Bkn[0], ldb, + (hipblasDoubleComplex *) &beta_p[0], + (hipblasDoubleComplex **)&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; + auto err = cublasZgemmBatched(gridblasHandle, + hOpA, + hOpB, + m,n,k, + (cuDoubleComplex *) &alpha_p[0], + (cuDoubleComplex **)&Amk[0], lda, + (cuDoubleComplex **)&Bkn[0], ldb, + (cuDoubleComplex *) &beta_p[0], + (cuDoubleComplex **)&Cmn[0], ldc, + batchCount); + GRID_ASSERT(err==CUBLAS_STATUS_SUCCESS); +#endif +#ifdef GRID_SYCL + 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, + (ComplexD *) &alpha_p[0], + (const ComplexD **)&Amk[0], (const int64_t *)&lda64, + (const ComplexD **)&Bkn[0], (const int64_t *)&ldb64, + (ComplexD *) &beta_p[0], + (ComplexD **)&Cmn[0], (const int64_t *)&ldc64, + (int64_t)1,&batchCount64,std::vector()); + synchronise(); +#endif +#if !defined(GRID_SYCL) && !defined(GRID_CUDA) && !defined(GRID_HIP) + // 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(ComplexD)*(m*k+k*n+m*n)*batchCount; + } + /////////////////////////////////////////////////////////////////////////// // Single precision real GEMM ///////////////////////////////////////////////////////////////////////////