/************************************************************************************* Grid physics library, www.github.com/paboyle/Grid Source file: ./tests/solver/Test_split_mixedprec_batched.cc Copyright (C) 2026 Author: Peter Boyle 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 using namespace std; using namespace Grid; const int NBATCH = 3; const RealD TOLERANCE = 1.0e-8; template 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 full(P,UGrid); for(int p=0;p void CheckSplitBatched(const std::string &name, LinearOperatorBase &Linop_d, LinearOperatorBase &Linop_f, GridBase *grid_f, std::vector &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 *split = Linop_f.SplitClone(layout); GRID_ASSERT(split != nullptr); int P = split->Partitions; std::vector full(P,grid_f); std::vector back(P,grid_f); std::vector Mfull(P,grid_f); for(int p=0;pFieldGrid); FieldF s_out(split->FieldGrid); Grid_split(full,s_in); Grid_unsplit(back,s_in); for(int p=0;pLinop->HermOp(s_in,s_out); Grid_unsplit(back,s_out); for(int p=0;p sol_ref(NBatch,src[0].Grid()); std::vector sol_split(NBatch,src[0].Grid()); for(int i=0;i 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 seeds4({1,2,3,4}); std::vector 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::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 src4_o(NBATCH,UrbGrid_d); std::vector src5_o(NBATCH,FrbGrid_d); std::vector src5(NBATCH,FGrid_d); LatticeFermionD tmp4(UGrid_d); LatticeFermionD tmp5(FGrid_d); for(int i=0;i Ld(Dd); SchurDiagMooeeOperator 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 Ld(Dd); SchurDiagMooeeOperator 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 Ld(Dd); SchurDiagMooeeOperator 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 Ld(Dd); SchurDiagMooeeOperator 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 Ld(Dd); SchurDiagMooeeOperator Lf(Df); CheckSplitBatched("Mobius SchurDiagMooee antiperiodic",Ld,Lf,FrbGrid_f,src5_o,layout); } { SchurDiagOneOperator Ld(Dd); SchurDiagOneOperator Lf(Df); CheckSplitBatched("Mobius SchurDiagOne antiperiodic",Ld,Lf,FrbGrid_f,src5_o,layout); } { SchurDiagTwoOperator Ld(Dd); SchurDiagTwoOperator Lf(Df); CheckSplitBatched("Mobius SchurDiagTwo antiperiodic",Ld,Lf,FrbGrid_f,src5_o,layout); } { MdagMLinearOperator Ld(Dd); MdagMLinearOperator 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 gamma(Ls); for(int s=0;s Ld(Dd); SchurDiagMooeeOperator 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 Lf(Df); SplitOperator *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(); }