mirror of
https://github.com/paboyle/Grid.git
synced 2026-08-15 23:19:37 +01:00
Batch blas test extension
This commit is contained in:
@@ -250,6 +250,84 @@ int main (int argc, char ** argv)
|
||||
GRID_ASSERT(rel < 1.0e-4);
|
||||
}
|
||||
|
||||
////////////////////////////////////////////////////////////////
|
||||
// T3 : ComplexD explicit-ld gemmBatched (the RecursiveSchurInverse
|
||||
// merge primitive): same split-K miniature as T2, double precision.
|
||||
// On device builds this is the FIRST exercise of hipblasZ/cublasZ
|
||||
// gemmBatched through the strided overload.
|
||||
////////////////////////////////////////////////////////////////
|
||||
{
|
||||
const int nrows = 8;
|
||||
const int nrhs = 4;
|
||||
const int K = 256;
|
||||
const int NK = 8; // split-K chunks
|
||||
const int Kc = K/NK;
|
||||
|
||||
std::vector<ComplexD> Ahost((uint64_t)K*nrows);
|
||||
std::vector<ComplexD> Xhost((uint64_t)K*nrhs);
|
||||
for(auto &z : Ahost) z = ComplexD(dist(rng),dist(rng));
|
||||
for(auto &z : Xhost) z = ComplexD(dist(rng),dist(rng));
|
||||
|
||||
deviceVector<ComplexD> Adev(Ahost.size());
|
||||
deviceVector<ComplexD> Xdev(Xhost.size());
|
||||
deviceVector<ComplexD> Yref((uint64_t)nrows*nrhs);
|
||||
deviceVector<ComplexD> Ypart((uint64_t)NK*nrows*nrhs);
|
||||
acceleratorCopyToDevice(&Ahost[0],&Adev[0],Ahost.size()*sizeof(ComplexD));
|
||||
acceleratorCopyToDevice(&Xhost[0],&Xdev[0],Xhost.size()*sizeof(ComplexD));
|
||||
|
||||
ComplexD one (1.0,0.0);
|
||||
ComplexD zero(0.0,0.0);
|
||||
|
||||
// (a) reference: one compact batch=1 call (compact lda == K for OP_T)
|
||||
{
|
||||
deviceVector<ComplexD*> Ap(1), Xp(1), Yp(1);
|
||||
std::vector<ComplexD*> h(1);
|
||||
h[0]=&Adev[0]; acceleratorCopyToDevice(&h[0],&Ap[0],sizeof(ComplexD*));
|
||||
h[0]=&Xdev[0]; acceleratorCopyToDevice(&h[0],&Xp[0],sizeof(ComplexD*));
|
||||
h[0]=&Yref[0]; acceleratorCopyToDevice(&h[0],&Yp[0],sizeof(ComplexD*));
|
||||
blas.gemmBatched(GridBLAS_OP_T,GridBLAS_OP_N,
|
||||
nrows,nrhs,K,
|
||||
one, Ap, Xp, zero, Yp);
|
||||
blas.synchronise();
|
||||
}
|
||||
|
||||
// (b) split-K: NK chunk-pointers into the SAME allocations, lda/ldb = K
|
||||
{
|
||||
deviceVector<ComplexD*> Ap(NK), Xp(NK), Yp(NK);
|
||||
std::vector<ComplexD*> h(NK);
|
||||
for(int j=0;j<NK;j++) h[j] = &Adev[(uint64_t)j*Kc]; // K-offset slice
|
||||
acceleratorCopyToDevice(&h[0],&Ap[0],NK*sizeof(ComplexD*));
|
||||
for(int j=0;j<NK;j++) h[j] = &Xdev[(uint64_t)j*Kc];
|
||||
acceleratorCopyToDevice(&h[0],&Xp[0],NK*sizeof(ComplexD*));
|
||||
for(int j=0;j<NK;j++) h[j] = &Ypart[(uint64_t)j*nrows*nrhs];
|
||||
acceleratorCopyToDevice(&h[0],&Yp[0],NK*sizeof(ComplexD*));
|
||||
blas.gemmBatched(GridBLAS_OP_T,GridBLAS_OP_N,
|
||||
nrows,nrhs,Kc,
|
||||
one, Ap, /*lda*/ K,
|
||||
Xp, /*ldb*/ K,
|
||||
zero,Yp, /*ldc*/ nrows);
|
||||
blas.synchronise();
|
||||
}
|
||||
|
||||
std::vector<ComplexD> Yref_h((uint64_t)nrows*nrhs);
|
||||
std::vector<ComplexD> Ypart_h((uint64_t)NK*nrows*nrhs);
|
||||
acceleratorCopyFromDevice(&Yref[0], &Yref_h[0], Yref_h.size()*sizeof(ComplexD));
|
||||
acceleratorCopyFromDevice(&Ypart[0],&Ypart_h[0],Ypart_h.size()*sizeof(ComplexD));
|
||||
|
||||
double maxdev = 0.0;
|
||||
double maxval = 0.0;
|
||||
for(int i=0;i<nrows*nrhs;i++){
|
||||
ComplexD sum(0.0,0.0);
|
||||
for(int j=0;j<NK;j++) sum = sum + Ypart_h[(uint64_t)j*nrows*nrhs + i];
|
||||
maxdev = std::max(maxdev,(double)abs(sum-Yref_h[i]));
|
||||
maxval = std::max(maxval,(double)abs(Yref_h[i]));
|
||||
}
|
||||
double rel = maxdev/maxval;
|
||||
std::cout << GridLogMessage << "T3 split-K strided gemmBatched ComplexD max rel dev vs compact = "
|
||||
<< rel << ( rel < 1.0e-13 ? " PASS" : " FAIL" ) << std::endl;
|
||||
GRID_ASSERT(rel < 1.0e-13);
|
||||
}
|
||||
|
||||
std::cout << GridLogMessage << "All batched-blas tests PASSED" << std::endl;
|
||||
|
||||
Grid_finalize();
|
||||
|
||||
Reference in New Issue
Block a user