mirror of
https://github.com/paboyle/Grid.git
synced 2026-10-04 06:58:05 +01:00
353 lines
14 KiB
C++
353 lines
14 KiB
C++
/*************************************************************************************
|
|
|
|
Grid physics library, www.github.com/paboyle/Grid
|
|
|
|
Source file: ./tests/solver/Test_split_mixedprec_batched.cc
|
|
|
|
Copyright (C) 2026
|
|
|
|
Author: Peter Boyle <paboyle@ph.ed.ac.uk>
|
|
|
|
This program is free software; you can redistribute it and/or modify
|
|
it under the terms of the GNU General Public License as published by
|
|
the Free Software Foundation; either version 2 of the License, or
|
|
(at your option) any later version.
|
|
|
|
This program is distributed in the hope that it will be useful,
|
|
but WITHOUT ANY WARRANTY; without even the implied warranty of
|
|
MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
|
|
GNU General Public License for more details.
|
|
|
|
You should have received a copy of the GNU General Public License along
|
|
with this program; if not, write to the Free Software Foundation, Inc.,
|
|
51 Franklin Street, Fifth Floor, Boston, MA 02110-1301 USA.
|
|
|
|
See the full license in the file "LICENSE" in the top level distribution directory
|
|
*************************************************************************************/
|
|
/* END LEGAL */
|
|
/////////////////////////////////////////////////////////////////////////////////////////////
|
|
// MixedPrecisionConjugateGradientBatched with inner solves on split-communicator
|
|
// partitions (--batched-solver-split), against the same solver unsplit.
|
|
//
|
|
// mpirun -n 2 ./Test_split_mixedprec_batched --grid 8.8.8.8 --mpi 1.1.1.2 --batched-solver-split 1.1.1.1
|
|
//
|
|
// With no --batched-solver-split the partitions are single ranks (1.1.1.1).
|
|
// NBATCH is deliberately not a multiple of the partition count, so the last group is
|
|
// zero-padded.
|
|
/////////////////////////////////////////////////////////////////////////////////////////////
|
|
#include <Grid/Grid.h>
|
|
|
|
using namespace std;
|
|
using namespace Grid;
|
|
|
|
const int NBATCH = 3;
|
|
const RealD TOLERANCE = 1.0e-8;
|
|
|
|
template<class Field>
|
|
RealD RelativeDifference(const Field &a,const Field &b)
|
|
{
|
|
Field diff(a.Grid());
|
|
diff = a - b;
|
|
return std::sqrt(norm2(diff)/norm2(b));
|
|
}
|
|
|
|
/////////////////////////////////////////////////////////////////////////////////////////////
|
|
// The solver assumes vector p of a group lands in the partition whose
|
|
// GridSplitVectorIndex is p: fill vector p with the constant p+1, split, and check
|
|
// every partition sees its own index.
|
|
/////////////////////////////////////////////////////////////////////////////////////////////
|
|
void CheckPartitionOrder(GridCartesian *UGrid,const Coordinate &layout)
|
|
{
|
|
GridCartesian SGrid(UGrid->FullDimensions(),UGrid->_simd_layout,layout,*UGrid);
|
|
int P = UGrid->ProcessorCount()/SGrid.ProcessorCount();
|
|
int index = GridSplitVectorIndex(UGrid,&SGrid);
|
|
|
|
std::vector<LatticeComplexF> full(P,UGrid);
|
|
for(int p=0;p<P;p++){
|
|
full[p] = ComplexF(p+1,0.0);
|
|
}
|
|
LatticeComplexF split(&SGrid);
|
|
Grid_split(full,split);
|
|
|
|
LatticeComplexF expect(&SGrid);
|
|
expect = ComplexF(index+1,0.0);
|
|
RealD err = norm2(split - expect);
|
|
|
|
std::cout << GridLogMessage << "Partition order: vector index " << index << " error " << err << std::endl;
|
|
GRID_ASSERT(err == 0.0);
|
|
}
|
|
|
|
/////////////////////////////////////////////////////////////////////////////////////////////
|
|
// For one pair of double/float linear operators:
|
|
// (1) red-black-aware split/unsplit round trip is exact;
|
|
// (2) the split clone's HermOp agrees with the full-grid HermOp;
|
|
// (3) the batched solver with split inner solves agrees with the unsplit solver;
|
|
// (4) the split solution has true residual at the requested tolerance.
|
|
/////////////////////////////////////////////////////////////////////////////////////////////
|
|
template<class FieldD,class FieldF>
|
|
void CheckSplitBatched(const std::string &name,
|
|
LinearOperatorBase<FieldD> &Linop_d,
|
|
LinearOperatorBase<FieldF> &Linop_f,
|
|
GridBase *grid_f,
|
|
std::vector<FieldD> &src,
|
|
const Coordinate &layout)
|
|
{
|
|
std::cout << GridLogMessage << "==================================================" << std::endl;
|
|
std::cout << GridLogMessage << name << std::endl;
|
|
std::cout << GridLogMessage << "==================================================" << std::endl;
|
|
|
|
int cb = src[0].Checkerboard();
|
|
|
|
SplitOperator<FieldF> *split = Linop_f.SplitClone(layout);
|
|
GRID_ASSERT(split != nullptr);
|
|
int P = split->Partitions;
|
|
|
|
std::vector<FieldF> full(P,grid_f);
|
|
std::vector<FieldF> back(P,grid_f);
|
|
std::vector<FieldF> Mfull(P,grid_f);
|
|
for(int p=0;p<P;p++){
|
|
full[p].Checkerboard() = cb;
|
|
back[p].Checkerboard() = cb;
|
|
Mfull[p].Checkerboard() = cb;
|
|
precisionChange(full[p],src[p%src.size()]);
|
|
}
|
|
|
|
// (1) round trip
|
|
FieldF s_in(split->FieldGrid);
|
|
FieldF s_out(split->FieldGrid);
|
|
Grid_split(full,s_in);
|
|
Grid_unsplit(back,s_in);
|
|
for(int p=0;p<P;p++){
|
|
RealD err = norm2(back[p] - full[p]);
|
|
std::cout << GridLogMessage << name << ": split/unsplit round trip " << p << " error " << err << std::endl;
|
|
GRID_ASSERT(err == 0.0);
|
|
}
|
|
|
|
// (2) operator
|
|
split->Linop->HermOp(s_in,s_out);
|
|
Grid_unsplit(back,s_out);
|
|
for(int p=0;p<P;p++){
|
|
Linop_f.HermOp(full[p],Mfull[p]);
|
|
RealD err = RelativeDifference(back[p],Mfull[p]);
|
|
std::cout << GridLogMessage << name << ": split HermOp vs full HermOp " << p << " relative difference " << err << std::endl;
|
|
GRID_ASSERT(err < 1.0e-6);
|
|
}
|
|
delete split;
|
|
|
|
// (3) solver, unsplit then split
|
|
int NBatch = src.size();
|
|
std::vector<FieldD> sol_ref(NBatch,src[0].Grid());
|
|
std::vector<FieldD> sol_split(NBatch,src[0].Grid());
|
|
for(int i=0;i<NBatch;i++){
|
|
sol_ref[i].Checkerboard() = cb;
|
|
sol_split[i].Checkerboard() = cb;
|
|
sol_ref[i] = Zero();
|
|
sol_split[i] = Zero();
|
|
}
|
|
|
|
MixedPrecisionConjugateGradientBatched<FieldD,FieldF> mCG(TOLERANCE,10000,50,1000,grid_f,Linop_f,Linop_d);
|
|
|
|
std::cout << GridLogMessage << name << ": unsplit batched solve" << std::endl;
|
|
mCG.BatchedSplit = Coordinate();
|
|
mCG.BatchedSplitNode = false;
|
|
mCG(src,sol_ref);
|
|
|
|
std::cout << GridLogMessage << name << ": split batched solve, partition layout " << layout << std::endl;
|
|
mCG.BatchedSplit = layout;
|
|
mCG(src,sol_split);
|
|
|
|
FieldD Msol(src[0].Grid());
|
|
Msol.Checkerboard() = cb;
|
|
for(int i=0;i<NBatch;i++){
|
|
RealD diff = RelativeDifference(sol_split[i],sol_ref[i]);
|
|
Linop_d.HermOp(sol_split[i],Msol);
|
|
RealD resid = RelativeDifference(Msol,src[i]);
|
|
std::cout << GridLogMessage << name << ": rhs " << i
|
|
<< " split vs unsplit " << diff
|
|
<< " true residual " << resid << std::endl;
|
|
GRID_ASSERT(diff < 1.0e-6);
|
|
GRID_ASSERT(resid < 10.0*TOLERANCE);
|
|
}
|
|
}
|
|
|
|
int main (int argc, char ** argv)
|
|
{
|
|
Grid_init(&argc,&argv);
|
|
|
|
const int Ls = 8;
|
|
|
|
Coordinate layout = GridDefaultBatchedSolverSplit();
|
|
if ( layout.size() == 0 ) {
|
|
layout = Coordinate(Nd,1);
|
|
}
|
|
|
|
GridCartesian *UGrid_d = SpaceTimeGrid::makeFourDimGrid(GridDefaultLatt(), GridDefaultSimd(Nd,vComplexD::Nsimd()), GridDefaultMpi());
|
|
GridRedBlackCartesian *UrbGrid_d = SpaceTimeGrid::makeFourDimRedBlackGrid(UGrid_d);
|
|
GridCartesian *FGrid_d = SpaceTimeGrid::makeFiveDimGrid(Ls,UGrid_d);
|
|
GridRedBlackCartesian *FrbGrid_d = SpaceTimeGrid::makeFiveDimRedBlackGrid(Ls,UGrid_d);
|
|
|
|
GridCartesian *UGrid_f = SpaceTimeGrid::makeFourDimGrid(GridDefaultLatt(), GridDefaultSimd(Nd,vComplexF::Nsimd()), GridDefaultMpi());
|
|
GridRedBlackCartesian *UrbGrid_f = SpaceTimeGrid::makeFourDimRedBlackGrid(UGrid_f);
|
|
GridCartesian *FGrid_f = SpaceTimeGrid::makeFiveDimGrid(Ls,UGrid_f);
|
|
GridRedBlackCartesian *FrbGrid_f = SpaceTimeGrid::makeFiveDimRedBlackGrid(Ls,UGrid_f);
|
|
|
|
CheckPartitionOrder(UGrid_f,layout);
|
|
|
|
std::vector<int> seeds4({1,2,3,4});
|
|
std::vector<int> seeds5({5,6,7,8});
|
|
GridParallelRNG RNG4(UGrid_d);
|
|
GridParallelRNG RNG5(FGrid_d);
|
|
RNG4.SeedFixedIntegers(seeds4);
|
|
RNG5.SeedFixedIntegers(seeds5);
|
|
|
|
LatticeGaugeFieldD Umu_d(UGrid_d);
|
|
LatticeGaugeFieldF Umu_f(UGrid_f);
|
|
SU<Nc>::HotConfiguration(RNG4,Umu_d);
|
|
precisionChange(Umu_f,Umu_d);
|
|
|
|
// Antiperiodic in time for some actions, so boundary phases must travel with the links
|
|
WilsonImplParams antiperiodic;
|
|
antiperiodic.boundary_phases[Nd-1] = -1.0;
|
|
|
|
// Sources: odd checkerboard (Schur) and full lattice (MdagM), 4d and 5d
|
|
std::vector<LatticeFermionD> src4_o(NBATCH,UrbGrid_d);
|
|
std::vector<LatticeFermionD> src5_o(NBATCH,FrbGrid_d);
|
|
std::vector<LatticeFermionD> src5(NBATCH,FGrid_d);
|
|
LatticeFermionD tmp4(UGrid_d);
|
|
LatticeFermionD tmp5(FGrid_d);
|
|
for(int i=0;i<NBATCH;i++){
|
|
random(RNG4,tmp4);
|
|
random(RNG5,tmp5);
|
|
pickCheckerboard(Odd,src4_o[i],tmp4);
|
|
pickCheckerboard(Odd,src5_o[i],tmp5);
|
|
src5[i] = tmp5;
|
|
}
|
|
|
|
//////////////////////////////////////////
|
|
// Wilson
|
|
//////////////////////////////////////////
|
|
{
|
|
RealD mass = 0.1;
|
|
WilsonFermionD Dd(Umu_d,*UGrid_d,*UrbGrid_d,mass);
|
|
WilsonFermionF Df(Umu_f,*UGrid_f,*UrbGrid_f,mass);
|
|
SchurDiagMooeeOperator<WilsonFermionD,LatticeFermionD> Ld(Dd);
|
|
SchurDiagMooeeOperator<WilsonFermionF,LatticeFermionF> Lf(Df);
|
|
CheckSplitBatched("Wilson SchurDiagMooee",Ld,Lf,UrbGrid_f,src4_o,layout);
|
|
}
|
|
|
|
//////////////////////////////////////////
|
|
// Wilson clover
|
|
//////////////////////////////////////////
|
|
{
|
|
RealD mass = 0.1;
|
|
RealD csw_r = 1.0;
|
|
RealD csw_t = 1.0;
|
|
WilsonCloverFermionD Dd(Umu_d,*UGrid_d,*UrbGrid_d,mass,csw_r,csw_t);
|
|
WilsonCloverFermionF Df(Umu_f,*UGrid_f,*UrbGrid_f,mass,csw_r,csw_t);
|
|
SchurDiagMooeeOperator<WilsonCloverFermionD,LatticeFermionD> Ld(Dd);
|
|
SchurDiagMooeeOperator<WilsonCloverFermionF,LatticeFermionF> Lf(Df);
|
|
CheckSplitBatched("WilsonClover SchurDiagMooee",Ld,Lf,UrbGrid_f,src4_o,layout);
|
|
}
|
|
|
|
//////////////////////////////////////////
|
|
// Compact Wilson clover, antiperiodic
|
|
//////////////////////////////////////////
|
|
{
|
|
RealD mass = 0.1;
|
|
RealD csw_r = 1.0;
|
|
RealD csw_t = 1.0;
|
|
RealD cF = 1.0;
|
|
WilsonAnisotropyCoefficients anis;
|
|
CompactWilsonCloverFermionD Dd(Umu_d,*UGrid_d,*UrbGrid_d,mass,csw_r,csw_t,cF,anis,antiperiodic);
|
|
CompactWilsonCloverFermionF Df(Umu_f,*UGrid_f,*UrbGrid_f,mass,csw_r,csw_t,cF,anis,antiperiodic);
|
|
SchurDiagMooeeOperator<CompactWilsonCloverFermionD,LatticeFermionD> Ld(Dd);
|
|
SchurDiagMooeeOperator<CompactWilsonCloverFermionF,LatticeFermionF> Lf(Df);
|
|
CheckSplitBatched("CompactWilsonClover SchurDiagMooee antiperiodic",Ld,Lf,UrbGrid_f,src4_o,layout);
|
|
}
|
|
|
|
//////////////////////////////////////////
|
|
// Domain wall
|
|
//////////////////////////////////////////
|
|
{
|
|
RealD mass = 0.1;
|
|
RealD M5 = 1.8;
|
|
DomainWallFermionD Dd(Umu_d,*FGrid_d,*FrbGrid_d,*UGrid_d,*UrbGrid_d,mass,M5);
|
|
DomainWallFermionF Df(Umu_f,*FGrid_f,*FrbGrid_f,*UGrid_f,*UrbGrid_f,mass,M5);
|
|
SchurDiagMooeeOperator<DomainWallFermionD,LatticeFermionD> Ld(Dd);
|
|
SchurDiagMooeeOperator<DomainWallFermionF,LatticeFermionF> Lf(Df);
|
|
CheckSplitBatched("DomainWall SchurDiagMooee",Ld,Lf,FrbGrid_f,src5_o,layout);
|
|
}
|
|
|
|
//////////////////////////////////////////
|
|
// Mobius, antiperiodic, with unequal
|
|
// masses; all wrapper kinds
|
|
//////////////////////////////////////////
|
|
{
|
|
RealD mass = 0.1;
|
|
RealD M5 = 1.8;
|
|
RealD b = 1.5;
|
|
RealD c = 0.5;
|
|
MobiusFermionD Dd(Umu_d,*FGrid_d,*FrbGrid_d,*UGrid_d,*UrbGrid_d,mass,M5,b,c,antiperiodic);
|
|
MobiusFermionF Df(Umu_f,*FGrid_f,*FrbGrid_f,*UGrid_f,*UrbGrid_f,mass,M5,b,c,antiperiodic);
|
|
Dd.SetMass(0.1,0.12);
|
|
Df.SetMass(0.1,0.12);
|
|
{
|
|
SchurDiagMooeeOperator<MobiusFermionD,LatticeFermionD> Ld(Dd);
|
|
SchurDiagMooeeOperator<MobiusFermionF,LatticeFermionF> Lf(Df);
|
|
CheckSplitBatched("Mobius SchurDiagMooee antiperiodic",Ld,Lf,FrbGrid_f,src5_o,layout);
|
|
}
|
|
{
|
|
SchurDiagOneOperator<MobiusFermionD,LatticeFermionD> Ld(Dd);
|
|
SchurDiagOneOperator<MobiusFermionF,LatticeFermionF> Lf(Df);
|
|
CheckSplitBatched("Mobius SchurDiagOne antiperiodic",Ld,Lf,FrbGrid_f,src5_o,layout);
|
|
}
|
|
{
|
|
SchurDiagTwoOperator<MobiusFermionD,LatticeFermionD> Ld(Dd);
|
|
SchurDiagTwoOperator<MobiusFermionF,LatticeFermionF> Lf(Df);
|
|
CheckSplitBatched("Mobius SchurDiagTwo antiperiodic",Ld,Lf,FrbGrid_f,src5_o,layout);
|
|
}
|
|
{
|
|
MdagMLinearOperator<MobiusFermionD,LatticeFermionD> Ld(Dd);
|
|
MdagMLinearOperator<MobiusFermionF,LatticeFermionF> Lf(Df);
|
|
CheckSplitBatched("Mobius MdagM antiperiodic",Ld,Lf,FGrid_f,src5,layout);
|
|
}
|
|
}
|
|
|
|
//////////////////////////////////////////
|
|
// ZMobius
|
|
//////////////////////////////////////////
|
|
{
|
|
RealD mass = 0.1;
|
|
RealD M5 = 1.8;
|
|
RealD b = 1.0;
|
|
RealD c = 0.0;
|
|
std::vector<ComplexD> gamma(Ls);
|
|
for(int s=0;s<Ls;s++){
|
|
gamma[s] = ComplexD(1.0+0.05*s, (s%2) ? 0.02 : -0.02);
|
|
}
|
|
ZMobiusFermionD Dd(Umu_d,*FGrid_d,*FrbGrid_d,*UGrid_d,*UrbGrid_d,mass,M5,gamma,b,c);
|
|
ZMobiusFermionF Df(Umu_f,*FGrid_f,*FrbGrid_f,*UGrid_f,*UrbGrid_f,mass,M5,gamma,b,c);
|
|
SchurDiagMooeeOperator<ZMobiusFermionD,LatticeFermionD> Ld(Dd);
|
|
SchurDiagMooeeOperator<ZMobiusFermionF,LatticeFermionF> Lf(Df);
|
|
CheckSplitBatched("ZMobius SchurDiagMooee",Ld,Lf,FrbGrid_f,src5_o,layout);
|
|
}
|
|
|
|
//////////////////////////////////////////
|
|
// A derived operator must not inherit
|
|
// its parent's SplitClone
|
|
//////////////////////////////////////////
|
|
{
|
|
RealD mass = 0.1;
|
|
RealD mu = 0.1;
|
|
WilsonTMFermionF Df(Umu_f,*UGrid_f,*UrbGrid_f,mass,mu);
|
|
SchurDiagMooeeOperator<WilsonTMFermionF,LatticeFermionF> Lf(Df);
|
|
SplitOperator<LatticeFermionF> *split = Lf.SplitClone(layout);
|
|
std::cout << GridLogMessage << "WilsonTM SplitClone refused: " << (split == nullptr) << std::endl;
|
|
GRID_ASSERT(split == nullptr);
|
|
}
|
|
|
|
std::cout << GridLogMessage << "Test_split_mixedprec_batched: all checks passed" << std::endl;
|
|
|
|
Grid_finalize();
|
|
}
|