Support for distributed Schur inverse

This commit is contained in:
Peter Boyle
2026-08-14 17:53:35 -04:00
parent 02d0301c9f
commit a7160ac513
+227
View File
@@ -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<ComplexD*> &Amk, int lda,
deviceVector<ComplexD*> &Bkn, int ldb,
ComplexD beta,
deviceVector<ComplexD*> &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<ComplexD> alpha_c;
static GridBLASDeviceConstant<ComplexD> 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<sycl::event>());
synchronise();
#endif
#if !defined(GRID_SYCL) && !defined(GRID_CUDA) && !defined(GRID_HIP)
// Reference implementation: Eigen with explicit outer stride
typedef Eigen::Map<Eigen::MatrixXcd,0,Eigen::OuterStride<> > 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
///////////////////////////////////////////////////////////////////////////