mirror of
https://github.com/paboyle/Grid.git
synced 2026-08-15 15:09:36 +01:00
Support for distributed Schur inverse
This commit is contained in:
@@ -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
|
||||
///////////////////////////////////////////////////////////////////////////
|
||||
|
||||
Reference in New Issue
Block a user