Block-Structured AMR Software Framework
Loading...
Searching...
No Matches
AMReX_MLCGSolver.H
Go to the documentation of this file.
1
2#ifndef AMREX_MLCGSOLVER_H_
3#define AMREX_MLCGSOLVER_H_
4#include <AMReX_Config.H>
5
6#include <AMReX_MLLinOp.H>
7
8namespace amrex {
9
19template <typename MF>
21{
22public:
23
24 using FAB = typename MLLinOpT<MF>::FAB;
25 using RT = typename MLLinOpT<MF>::RT;
26
27 enum struct Type { BiCGStab, CG };
28
37
38 MLCGSolverT (const MLCGSolverT<MF>& rhs) = delete;
39 MLCGSolverT (MLCGSolverT<MF>&& rhs) = delete;
42
48 void setSolver (Type _typ) noexcept { solver_type = _typ; }
49
65 int solve (MF& solnL, const MF& rhsL, RT eps_rel, RT eps_abs);
66
72 void setVerbose (int _verbose) { verbose = _verbose; }
74 [[nodiscard]] int getVerbose () const { return verbose; }
75
81 void setMaxIter (int _maxiter) { maxiter = _maxiter; }
83 [[nodiscard]] int getMaxIter () const { return maxiter; }
84
90 void setPrintIdentation (std::string s) { print_ident = std::move(s); }
91
99 void setInitSolnZeroed (bool _sol_zeroed) { initial_vec_zeroed = _sol_zeroed; }
101 [[nodiscard]] bool getInitSolnZeroed () const { return initial_vec_zeroed; }
102
108 void setNGhost(int _nghost) {nghost = IntVect(_nghost);}
110 [[nodiscard]] int getNGhost() {return nghost[0];}
111
120 [[nodiscard]] RT dotxy (const MF& r, const MF& z, bool local = false);
128 [[nodiscard]] RT norm_inf (const MF& res, bool local = false);
138 int solve_bicgstab (MF& solnL, const MF& rhsL, RT eps_rel, RT eps_abs);
148 int solve_cg (MF& solnL, const MF& rhsL, RT eps_rel, RT eps_abs);
149
151 [[nodiscard]] int getNumIters () const noexcept { return iter; }
152
153private:
154
155 MLLinOpT<MF>& Lp;
156 Type solver_type;
157 const int amrlev = 0;
158 const int mglev;
159 int verbose = 0;
160 int maxiter = 100;
161 IntVect nghost = IntVect(0);
162 int iter = -1;
163 bool initial_vec_zeroed = false;
164 std::string print_ident;
165};
166
167template <typename MF>
169 : Lp(_lp), solver_type(_typ), mglev(_lp.NMGLevels(0)-1)
170{}
171
172template <typename MF> MLCGSolverT<MF>::~MLCGSolverT () = default;
173
174template <typename MF>
175int
176MLCGSolverT<MF>::solve (MF& sol, const MF& rhs, RT eps_rel, RT eps_abs)
177{
178 if (solver_type == Type::BiCGStab) {
179 return solve_bicgstab(sol,rhs,eps_rel,eps_abs);
180 } else {
181 return solve_cg(sol,rhs,eps_rel,eps_abs);
182 }
183}
184
185template <typename MF>
186int
187MLCGSolverT<MF>::solve_bicgstab (MF& sol, const MF& rhs, RT eps_rel, RT eps_abs)
188{
189 BL_PROFILE("MLCGSolver::bicgstab");
190
191 const int ncomp = nComp(sol);
192
193 MFInfo const async_info = MFInfo().SetArena(The_Async_Arena());
194 MF p = Lp.make(amrlev, mglev, nGrowVect(sol), async_info);
195 MF r = Lp.make(amrlev, mglev, nGrowVect(sol), async_info);
196 setVal(p, RT(0.0)); // Make sure all entries are initialized to avoid errors
197 setVal(r, RT(0.0));
198
199 MF rh = Lp.make(amrlev, mglev, nghost, async_info);
200 MF v = Lp.make(amrlev, mglev, nghost, async_info);
201 MF t = Lp.make(amrlev, mglev, nghost, async_info);
202
203
204 MF sorig;
205
206 if ( initial_vec_zeroed ) {
207 LocalCopy(r,rhs,0,0,ncomp,nghost);
208 } else {
209 sorig = Lp.make(amrlev, mglev, nghost, async_info);
210
211 Lp.correctionResidual(amrlev, mglev, r, sol, rhs, MLLinOpT<MF>::BCMode::Homogeneous);
212
213 LocalCopy(sorig,sol,0,0,ncomp,nghost);
214 setVal(sol, RT(0.0));
215 }
216
217 // Then normalize
218 Lp.normalize(amrlev, mglev, r);
219 LocalCopy(rh, r, 0,0,ncomp,nghost);
220
221 RT rnorm = norm_inf(r);
222 const RT rnorm0 = rnorm;
223
224 if ( verbose > 0 )
225 {
226 amrex::Print() << print_ident << "MLCGSolver_BiCGStab: Initial error (error0) = " << rnorm0 << '\n';
227 }
228 int ret = 0;
229 iter = 1;
230 RT rho_1 = 0, alpha = 0, omega = 0;
231
232 if ( rnorm0 == 0 || rnorm0 < eps_abs )
233 {
234 if ( verbose > 0 )
235 {
236 amrex::Print() << print_ident << "MLCGSolver_BiCGStab: niter = 0,"
237 << ", rnorm = " << rnorm
238 << ", eps_abs = " << eps_abs << '\n';
239 }
240 if ( !initial_vec_zeroed ) {
241 LocalAdd(sol, sorig, 0, 0, ncomp, nghost);
242 }
243 return ret;
244 }
245
246 for (; iter <= maxiter; ++iter)
247 {
248 const RT rho = dotxy(rh,r);
249 if ( rho == 0 )
250 {
251 ret = 1; break;
252 }
253 if ( iter == 1 )
254 {
255 LocalCopy(p,r,0,0,ncomp,nghost);
256 }
257 else
258 {
259 const RT beta = (rho/rho_1)*(alpha/omega);
260 if constexpr (IsMultiFabLike_v<MF>) {
261 // two operations: p += -omega*v; p = r + beta*p
262 // same as: p = r + beta*(p - omega*v)
263 Saxpy_Xpay(p, -omega, v, beta, r, 0, 0, ncomp, nghost);
264 } else {
265 Saxpy(p, -omega, v, 0, 0, ncomp, nghost); // p += -omega*v
266 Xpay(p, beta, r, 0, 0, ncomp, nghost); // p = r + beta*p
267 }
268 }
270 Lp.normalize(amrlev, mglev, v);
271
272 RT rhTv = dotxy(rh,v);
273 if ( rhTv != RT(0.0) )
274 {
275 alpha = rho/rhTv;
276 }
277 else
278 {
279 ret = 2; break;
280 }
281 if constexpr (IsMultiFabLike_v<MF>) {
282 // sol += alpha * p; r += -alpha * v
283 Saxpy_Saxpy(sol, alpha, p, r, -alpha, v, 0, 0, ncomp, nghost);
284 } else {
285 Saxpy(sol, alpha, p, 0, 0, ncomp, nghost); // sol += alpha * p
286 Saxpy(r, -alpha, v, 0, 0, ncomp, nghost); // r += -alpha * v
287 }
288
289 rnorm = norm_inf(r);
290
291 if ( verbose > 2 && ParallelDescriptor::IOProcessor() )
292 {
293 amrex::Print() << print_ident << "MLCGSolver_BiCGStab: Half Iter "
294 << std::setw(11) << iter
295 << " rel. err. "
296 << rnorm/(rnorm0) << '\n';
297 }
298
299 if ( rnorm < eps_rel*rnorm0 || rnorm < eps_abs ) { break; }
300
302 Lp.normalize(amrlev, mglev, t);
303 //
304 // This is a little funky. I want to elide one of the reductions
305 // in the following two dotxy()s. We do that by calculating the "local"
306 // values and then reducing the two local values at the same time.
307 //
308 RT tvals[2] = { dotxy(t,t,true), dotxy(t,r,true) };
309
310 BL_PROFILE_VAR("MLCGSolver::ParallelAllReduce", blp_par);
311 ParallelAllReduce::Sum(tvals,2,Lp.BottomCommunicator());
312 BL_PROFILE_VAR_STOP(blp_par);
313
314 if ( tvals[0] != RT(0.0) )
315 {
316 omega = tvals[1]/tvals[0];
317 }
318 else
319 {
320 ret = 3; break;
321 }
322 if constexpr (IsMultiFabLike_v<MF>) {
323 // sol += omega * r; r += -omega * t
324 Saypy_Saxpy(sol, omega, r, -omega, t, 0, 0, ncomp, nghost);
325 } else {
326 Saxpy(sol, omega, r, 0, 0, ncomp, nghost); // sol += omega * r
327 Saxpy(r, -omega, t, 0, 0, ncomp, nghost); // r += -omega * t
328 }
329
330 rnorm = norm_inf(r);
331
332 if ( verbose > 2 )
333 {
334 amrex::Print() << print_ident << "MLCGSolver_BiCGStab: Iteration "
335 << std::setw(11) << iter
336 << " rel. err. "
337 << rnorm/(rnorm0) << '\n';
338 }
339
340 if ( rnorm < eps_rel*rnorm0 || rnorm < eps_abs ) { break; }
341
342 if ( omega == 0 )
343 {
344 ret = 4; break;
345 }
346 rho_1 = rho;
347 }
348
349 if ( verbose > 0 )
350 {
351 amrex::Print() << print_ident << "MLCGSolver_BiCGStab: Final: Iteration "
352 << std::setw(4) << iter
353 << " rel. err. "
354 << rnorm/(rnorm0) << '\n';
355 }
356
357 if ( ret == 0 && rnorm > eps_rel*rnorm0 && rnorm > eps_abs)
358 {
359 if ( verbose > 0 && ParallelDescriptor::IOProcessor() ) {
360 amrex::Warning("MLCGSolver_BiCGStab:: failed to converge!");
361 }
362 ret = 8;
363 }
364
365 if ( ( ret == 0 || ret == 8 ) && (rnorm < rnorm0) )
366 {
367 if ( !initial_vec_zeroed ) {
368 LocalAdd(sol, sorig, 0, 0, ncomp, nghost);
369 }
370 if (ret == 8) { ret = 9; }
371 }
372 else
373 {
374 setVal(sol, RT(0.0));
375 if ( !initial_vec_zeroed ) {
376 LocalAdd(sol, sorig, 0, 0, ncomp, nghost);
377 }
378 }
379
380 return ret;
381}
382
383template <typename MF>
384int
385MLCGSolverT<MF>::solve_cg (MF& sol, const MF& rhs, RT eps_rel, RT eps_abs)
386{
387 BL_PROFILE("MLCGSolver::cg");
388
389 const int ncomp = nComp(sol);
390
391 MFInfo const async_info = MFInfo().SetArena(The_Async_Arena());
392 MF p = Lp.make(amrlev, mglev, nGrowVect(sol), async_info);
393 setVal(p, RT(0.0));
394
395 MF r = Lp.make(amrlev, mglev, nghost, async_info);
396 MF q = Lp.make(amrlev, mglev, nghost, async_info);
397
398 MF sorig;
399
400 if ( initial_vec_zeroed ) {
401 LocalCopy(r,rhs,0,0,ncomp,nghost);
402 } else {
403 sorig = Lp.make(amrlev, mglev, nghost, async_info);
404
405 Lp.correctionResidual(amrlev, mglev, r, sol, rhs, MLLinOpT<MF>::BCMode::Homogeneous);
406
407 LocalCopy(sorig,sol,0,0,ncomp,nghost);
408 setVal(sol, RT(0.0));
409 }
410
411 RT rnorm = norm_inf(r);
412 const RT rnorm0 = rnorm;
413
414 if ( verbose > 0 )
415 {
416 amrex::Print() << print_ident << "MLCGSolver_CG: Initial error (error0) : " << rnorm0 << '\n';
417 }
418
419 RT rho_1 = 0;
420 int ret = 0;
421 iter = 1;
422
423 if ( rnorm0 == 0 || rnorm0 < eps_abs )
424 {
425 if ( verbose > 0 ) {
426 amrex::Print() << print_ident << "MLCGSolver_CG: niter = 0,"
427 << ", rnorm = " << rnorm
428 << ", eps_abs = " << eps_abs << '\n';
429 }
430 if ( !initial_vec_zeroed ) {
431 LocalAdd(sol, sorig, 0, 0, ncomp, nghost);
432 }
433 return ret;
434 }
435
436 for (; iter <= maxiter; ++iter)
437 {
438 RT rho = dotxy(r,r);
439
440 if ( rho == 0 )
441 {
442 ret = 1; break;
443 }
444 if (iter == 1)
445 {
446 LocalCopy(p,r,0,0,ncomp,nghost);
447 }
448 else
449 {
450 RT beta = rho/rho_1;
451 Xpay(p, beta, r, 0, 0, ncomp, nghost); // p = r + beta * p
452 }
454
455 RT alpha;
456 RT pw = dotxy(p,q);
457 if ( pw != RT(0.0))
458 {
459 alpha = rho/pw;
460 }
461 else
462 {
463 ret = 1; break;
464 }
465
466 if ( verbose > 2 )
467 {
468 amrex::Print() << print_ident << "MLCGSolver_cg:"
469 << " iter " << iter
470 << " rho " << rho
471 << " alpha " << alpha << '\n';
472 }
473 if constexpr (IsMultiFabLike_v<MF>) {
474 // sol += alpha * p; r += -alpha * q
475 Saxpy_Saxpy(sol, alpha, p, r, -alpha, q, 0, 0, ncomp, nghost);
476 } else {
477 Saxpy(sol, alpha, p, 0, 0, ncomp, nghost); // sol += alpha * p
478 Saxpy(r, -alpha, q, 0, 0, ncomp, nghost); // r += -alpha * q
479 }
480 rnorm = norm_inf(r);
481
482 if ( verbose > 2 )
483 {
484 amrex::Print() << print_ident << "MLCGSolver_cg: Iteration"
485 << std::setw(4) << iter
486 << " rel. err. "
487 << rnorm/(rnorm0) << '\n';
488 }
489
490 if ( rnorm < eps_rel*rnorm0 || rnorm < eps_abs ) { break; }
491
492 rho_1 = rho;
493 }
494
495 if ( verbose > 0 )
496 {
497 amrex::Print() << print_ident << "MLCGSolver_cg: Final Iteration"
498 << std::setw(4) << iter
499 << " rel. err. "
500 << rnorm/(rnorm0) << '\n';
501 }
502
503 if ( ret == 0 && rnorm > eps_rel*rnorm0 && rnorm > eps_abs )
504 {
505 if ( verbose > 0 && ParallelDescriptor::IOProcessor() ) {
506 amrex::Warning("MLCGSolver_cg: failed to converge!");
507 }
508 ret = 8;
509 }
510
511 if ( ( ret == 0 || ret == 8 ) && (rnorm < rnorm0) )
512 {
513 if ( !initial_vec_zeroed ) {
514 LocalAdd(sol, sorig, 0, 0, ncomp, nghost);
515 }
516 if (ret == 8) { ret = 9; }
517 }
518 else
519 {
520 setVal(sol, RT(0.0));
521 if ( !initial_vec_zeroed ) {
522 LocalAdd(sol, sorig, 0, 0, ncomp, nghost);
523 }
524 }
525
526 return ret;
527}
528
529template <typename MF>
530auto
531MLCGSolverT<MF>::dotxy (const MF& r, const MF& z, bool local) -> RT
532{
533 BL_PROFILE_VAR_NS("MLCGSolver::ParallelAllReduce", blp_par);
534 if (!local) { BL_PROFILE_VAR_START(blp_par); }
535 RT result = Lp.xdoty(amrlev, mglev, r, z, local);
536 if (!local) { BL_PROFILE_VAR_STOP(blp_par); }
537 return result;
538}
539
540template <typename MF>
541auto
542MLCGSolverT<MF>::norm_inf (const MF& res, bool local) -> RT
543{
544 int ncomp = nComp(res);
545 RT result = norminf(res,0,ncomp,IntVect(0),true);
546 if (!local) {
547 BL_PROFILE("MLCGSolver::ParallelAllReduce");
548 ParallelAllReduce::Max(result, Lp.BottomCommunicator());
549 }
550 return result;
551}
552
554
555}
556
557#endif /*_CGSOLVER_H_*/
#define BL_PROFILE_VAR_START(vname)
Definition AMReX_BLProfiler.H:573
#define BL_PROFILE(a)
Definition AMReX_BLProfiler.H:562
#define BL_PROFILE_VAR_STOP(vname)
Definition AMReX_BLProfiler.H:574
#define BL_PROFILE_VAR(fname, vname)
Definition AMReX_BLProfiler.H:571
#define BL_PROFILE_VAR_NS(fname, vname)
Definition AMReX_BLProfiler.H:572
GpuArray< Real, 3 > beta
Definition AMReX_MLEBNodeFDLaplacian.cpp:1099
CG-family solvers (BiCGStab or CG) for use as the bottom solver in MLMG.
Definition AMReX_MLCGSolver.H:21
RT dotxy(const MF &r, const MF &z, bool local=false)
Dot product helper; set local to true to skip the MPI reduction.
Definition AMReX_MLCGSolver.H:531
void setSolver(Type _typ) noexcept
Switch between BiCGStab and CG after construction.
Definition AMReX_MLCGSolver.H:48
bool getInitSolnZeroed() const
Whether setInitSolnZeroed(true) was requested.
Definition AMReX_MLCGSolver.H:101
void setVerbose(int _verbose)
Control how much logging is emitted (0 = silent).
Definition AMReX_MLCGSolver.H:72
int getNGhost()
Current grow-cell count (same in every direction).
Definition AMReX_MLCGSolver.H:110
MLCGSolverT< MF > & operator=(const MLCGSolverT< MF > &rhs)=delete
int getNumIters() const noexcept
Iteration count from the last solve* call (or -1 if unused).
Definition AMReX_MLCGSolver.H:151
int solve_cg(MF &solnL, const MF &rhsL, RT eps_rel, RT eps_abs)
Raw CG implementation mirroring the high-level solve() signature.
Definition AMReX_MLCGSolver.H:385
RT norm_inf(const MF &res, bool local=false)
Infinity norm helper; set local to true to skip the MPI reduction.
Definition AMReX_MLCGSolver.H:542
void setInitSolnZeroed(bool _sol_zeroed)
Definition AMReX_MLCGSolver.H:99
int solve_bicgstab(MF &solnL, const MF &rhsL, RT eps_rel, RT eps_abs)
Raw BiCGStab implementation mirroring the high-level solve() signature.
Definition AMReX_MLCGSolver.H:187
MLCGSolverT(MLCGSolverT< MF > &&rhs)=delete
typename MLLinOpT< MF >::FAB FAB
Definition AMReX_MLCGSolver.H:24
int getMaxIter() const
Current iteration cap.
Definition AMReX_MLCGSolver.H:83
void setPrintIdentation(std::string s)
Prefix printed messages (e.g., to indent per level).
Definition AMReX_MLCGSolver.H:90
int solve(MF &solnL, const MF &rhsL, RT eps_rel, RT eps_abs)
Solve Lp(solnL)=rhsL to the requested tolerance.
Definition AMReX_MLCGSolver.H:176
Type
Definition AMReX_MLCGSolver.H:27
typename MLLinOpT< MF >::RT RT
Definition AMReX_MLCGSolver.H:25
void setNGhost(int _nghost)
Set the number of grow cells used when allocating temporaries.
Definition AMReX_MLCGSolver.H:108
MLCGSolverT(MLLinOpT< MF > &_lp, Type _typ=Type::BiCGStab)
Construct a solver bound to _lp.
Definition AMReX_MLCGSolver.H:168
int getVerbose() const
Current verbosity level.
Definition AMReX_MLCGSolver.H:74
void setMaxIter(int _maxiter)
Cap the number of Krylov iterations performed.
Definition AMReX_MLCGSolver.H:81
MLCGSolverT(const MLCGSolverT< MF > &rhs)=delete
Abstract base class for multilevel linear operators used by MLMG and the bottom solvers.
Definition AMReX_MLLinOp.H:137
typename FabDataType< MF >::fab_type FAB
Definition AMReX_MLLinOp.H:147
typename FabDataType< MF >::value_type RT
Definition AMReX_MLLinOp.H:148
This class provides the user with a few print options.
Definition AMReX_Print.H:35
Arena * The_Async_Arena()
Definition AMReX_Arena.cpp:825
void Sum(Gpu::DeviceVector< T > &v, MPI_Comm comm)
Definition AMReX_GpuParallelReduce.H:37
bool IOProcessor() noexcept
Is this CPU the I/O Processor? To get the rank number, call IOProcessorNumber()
Definition AMReX_ParallelDescriptor.H:289
void Max(KeyValuePair< K, V > &vi, MPI_Comm comm)
Definition AMReX_ParallelReduce.H:133
Definition AMReX_Amr.cpp:50
void Saxpy_Xpay(MF &dst, typename MF::value_type a_saxpy, MF const &src_saxpy, typename MF::value_type a_xpay, MF const &src_xpay, int scomp, int dcomp, int ncomp, IntVect const &nghost)
dst += a_saxpy * src_saxpy followed by dst = src_xpay + a_xpay * dst
Definition AMReX_FabArrayUtility.H:2214
int nComp(FabArrayBase const &fa)
Convenience wrapper that forwards to fa.nComp().
Definition AMReX_FabArrayBase.cpp:2860
void Saxpy_Saxpy(MF &dst1, typename MF::value_type a1, MF const &src1, MF &dst2, typename MF::value_type a2, MF const &src2, int scomp, int dcomp, int ncomp, IntVect const &nghost)
dst1 += a1 * src1 followed by dst2 += a2 * src2
Definition AMReX_FabArrayUtility.H:2223
void Saypy_Saxpy(MF &dst1, typename MF::value_type a1, MF &dst2, typename MF::value_type a2, MF const &src, int scomp, int dcomp, int ncomp, IntVect const &nghost)
dst1 += a1 * dst2 followed by dst2 += a2 * src
Definition AMReX_FabArrayUtility.H:2232
IntVectND< 3 > IntVect
IntVect is an alias for amrex::IntVectND instantiated with AMREX_SPACEDIM.
Definition AMReX_BaseFwd.H:38
IntVect nGrowVect(FabArrayBase const &fa)
Convenience wrapper that forwards to fa.nGrowVect().
Definition AMReX_FabArrayBase.cpp:2865
void LocalCopy(DMF &dst, SMF const &src, int scomp, int dcomp, int ncomp, IntVect const &nghost)
dst = src
Definition AMReX_FabArrayUtility.H:2182
MF::value_type norminf(MF const &mf, int scomp, int ncomp, IntVect const &nghost, bool local=false)
Return the infinity norm, with an MPI maximum unless local is true.
Definition AMReX_FabArrayUtility.H:2262
void Xpay(MF &dst, typename MF::value_type a, MF const &src, int scomp, int dcomp, int ncomp, IntVect const &nghost)
dst = src + a * dst
Definition AMReX_FabArrayUtility.H:2206
void Warning(const std::string &msg)
Print a warning message to the diagnostic stream and keep running.
Definition AMReX.cpp:248
void Saxpy(MF &dst, typename MF::value_type a, MF const &src, int scomp, int dcomp, int ncomp, IntVect const &nghost)
dst += a * src
Definition AMReX_FabArrayUtility.H:2198
void LocalAdd(MF &dst, MF const &src, int scomp, int dcomp, int ncomp, IntVect const &nghost)
dst += src
Definition AMReX_FabArrayUtility.H:2190
void setVal(MF &dst, typename MF::value_type val)
dst = val
Definition AMReX_FabArrayUtility.H:2161
FabArray memory allocation information.
Definition AMReX_FabArray.H:73
MFInfo & SetArena(Arena *ar) noexcept
Select the Arena used for FAB storage.
Definition AMReX_FabArray.H:87