"""Stokes, Taylor-Hood Q2/Q1: block preconditioners, iterations flat in h. -nu Delta u + grad p = f, -div u = 0, u = 0 on the boundary The same manufactured problem as ``solve_stokes_saddle_2d.py``; what changes is how the saddle point is solved. Unpreconditioned, the Krylov count grows like the number of unknowns (85, 580, 1336 BiCGStab iterations at 187, 659, 2467 DOFs, see ``EllipticSaddleScheme``). The textbook remedy (Elman, Silvester & Wathen, *Finite Elements and Fast Iterative Solvers*, ch. 4): precondition block by block, the velocity block ``A = nu K`` by a multigrid V-cycle, the Schur complement ``S = -B A^-1 B^T`` by the pressure mass, ``S ~ -M_p / nu`` -- spectrally equivalent, whatever h. The count then stays flat. Three forms, the SAME ``BlockSchurPreconditioner`` with the same two inverses, only the sweep order changing (see ``solvers/schur.py``): * block-diagonal ``diag(A, +M_p/nu)``, ``mode="jacobi"``: symmetric positive definite, so MINRES; * block-triangular ``[[A, B^T], [0, -M_p/nu]]``, ``sweep=[p, u]``: FGMRES; * block LU, ``sweep=[u, p, u]``: FGMRES, one more V-cycle per iteration. ``A^-1`` is ONE V-cycle on the velocity space (``MultigridInverse``: its levels' operator is read off the Stokes system's own element matrices, the coarse ones by Galerkin -- built once, Stokes being linear); ``M_p^-1`` is eight Jacobi-Chebyshev steps (``ChebyshevInverse``, interval by Wathen's element bound, [1/4, 9/4] for Q1). The pressure is determined up to a constant here (u prescribed all around): the system is singular but consistent, and both Krylov methods converge on it without any projection; the errors are measured on the mean-free pressure. Printed: Krylov iterations; the setup (hierarchy, multigrid, pressure mass) plus the FIRST solve, which compiles; the warm solve; then the DOFs and the L2 errors, which must not depend on the preconditioner. """ import time import jax import jax.numpy as jnp import numpy as np from scimba_jax.linear_approximation.basis.analytic_bases import local_lagrange_basis from scimba_jax.linear_approximation.basis.dof_map import LagrangeDofMap from scimba_jax.linear_approximation.basis.general_bases import AnalyticBasis from scimba_jax.linear_approximation.galerkin.fem.elliptic_fe_scheme import ( EllipticFEscheme, ) from scimba_jax.linear_approximation.galerkin.fem.elliptic_saddle_scheme import ( EllipticSaddleScheme, ) from scimba_jax.linear_approximation.meshes.mesh import Mesh from scimba_jax.linear_approximation.quad.gauss_quad import UnitSquareTensorized from scimba_jax.linear_approximation.solvers import LinearKrylov from scimba_jax.linear_approximation.solvers.multigrid import MG from scimba_jax.linear_approximation.solvers.schur import ( BlockSchurPreconditioner, ChebyshevInverse, MultigridInverse, ScaledInverse, ) from scimba_jax.linear_approximation.solvers.smoothers import DampedJacobiSmoother from scimba_jax.linear_approximation.transfer.hierarchy import ( build_hierarchy_structured, ) from scimba_jax.linear_approximation.variables.variables_fe import VariablesFE from scimba_jax.mapping.mapping import InvertibleFunction, Mapping from scimba_jax.physical_models.abstract_physical_weak_model import ( AbstractPhysicalWeakModel, ) from scimba_jax.physical_models.classical_weakform.mass_weak_form import ( MassWeakForm, ) from scimba_jax.physical_models.classical_weakform.stokes_weak_form import ( StokesWeakForm, VectorLaplacianWeakForm, ) from scimba_jax.physical_models.weak_boundary_conditions import Dirichlet DIM = 2 QUAD_ORDER = 6 NU = 1.0 N_CELLS = (8, 16, 32, 64) N_COARSE = 4 # the multigrid's coarsest grid TOL = 1e-8 U, P = 0, 1 def u_exact(x): """Divergence-free and vanishing on the whole boundary.""" sx, cx = jnp.sin(jnp.pi * x[0]), jnp.cos(jnp.pi * x[0]) sy, cy = jnp.sin(jnp.pi * x[1]), jnp.cos(jnp.pi * x[1]) return jnp.pi * jnp.array([sx * sx * sy * cy, -sy * sy * sx * cx]) def p_exact(x): """Zero mean on the square, since the pressure is fixed only up to one.""" return jnp.array([jnp.cos(jnp.pi * x[0]) * jnp.sin(jnp.pi * x[1])]) def source(x): """f = -nu Delta u + grad p, differentiated rather than written out.""" laplacian = jnp.trace(jax.jacfwd(jax.jacfwd(u_exact))(x), axis1=1, axis2=2) return -NU * laplacian + jax.jacfwd(lambda y: p_exact(y)[0])(x) def make_mesh(n_cells): return Mesh( dim=DIM, n_cells=(n_cells, n_cells), ref_quad=UnitSquareTensorized(dim=DIM, order=QUAD_ORDER), mapping=Mapping(mappings=[InvertibleFunction(lambda x: x, lambda y: y)]), ) def make_space(mesh, order, out_dim, basis_type): basis = AnalyticBasis( nb_basis=(order + 1) ** DIM, out_dim=out_dim, mesh=mesh, local_basis=lambda y, i, m: local_lagrange_basis( y, i, m, order=order, out_dim=out_dim ), basis_type=basis_type, ) return VariablesFE(basis=basis, nb_variables=out_dim, dof_map=LagrangeDofMap) def stokes_scheme(mesh): model = AbstractPhysicalWeakModel(dim=DIM) model.add_weak_form("main", StokesWeakForm(dim=DIM, f=source, viscosity=NU)) model.add_boundary_condition("0/boundary", Dirichlet(lambda x: jnp.zeros(DIM))) velocity = make_space(mesh, 2, DIM, "field") pressure = make_space(mesh, 1, 1, "scalar") return EllipticSaddleScheme(model, [velocity, pressure]) def velocity_scheme(n_cells): """The velocity space alone, for the multigrid's STRUCTURE (see the module).""" model = AbstractPhysicalWeakModel.from_weak_form( VectorLaplacianWeakForm(dim=DIM, viscosity=NU), dirichlet=lambda x: jnp.zeros(DIM), ) return EllipticFEscheme(model, make_space(make_mesh(n_cells[0]), 2, DIM, "field")) def pressure_mass(mesh): model = AbstractPhysicalWeakModel.from_weak_form(MassWeakForm(dim=DIM)) return EllipticFEscheme(model, make_space(mesh, 1, 1, "scalar")) def inverses(scheme, n_cells, schur_sign): """``A^-1`` by one V-cycle, ``S^-1`` by ``schur_sign * nu * M_p^-1``.""" n_levels = int(np.log2(n_cells // N_COARSE)) + 1 hierarchy = build_hierarchy_structured( velocity_scheme, (n_cells, n_cells), n_levels ) mg = MG(hierarchy, DampedJacobiSmoother(omega=2.0 / 3.0), nu_pre=2, nu_post=2) zeros = scheme._initial_dofs() velocity = MultigridInverse.frozen_component(mg, scheme, zeros, U) mass = ChebyshevInverse(pressure_mass(scheme.variables_list[P].mesh), degree=8) return {U: velocity, P: ScaledInverse(mass, schur_sign * NU)} def errors(dofsl, scheme): """L2 errors of u and of the mean-free p, on a 60x60 cell-centred grid.""" u_dofs, p_dofs = scheme._split_dofs(dofsl) grid = (np.arange(60) + 0.5) / 60 points = jnp.asarray( np.stack(np.meshgrid(grid, grid, indexing="ij"), axis=-1).reshape(-1, 2) ) velocity, pressure = scheme.variables_list velocity.dofsl, pressure.dofsl = u_dofs, p_dofs u_h, p_h = velocity.evaluate(points), pressure.evaluate(points)[:, 0] u_ref, p_ref = jax.vmap(u_exact)(points), jax.vmap(p_exact)(points)[:, 0] error_u = float(jnp.sqrt(jnp.mean(jnp.sum((u_h - u_ref) ** 2, axis=-1)))) error_p = float( jnp.sqrt(jnp.mean(((p_h - p_h.mean()) - (p_ref - p_ref.mean())) ** 2)) ) return error_u, error_p def solve(scheme, solver): """``(dofsl, report, first, warm)``: the first call compiles, the second not. Both are reported: on these sizes the compilation is most of what a single solve costs, and a table of warm times alone would hide it. """ started = time.perf_counter() EllipticSaddleScheme.solve_pure(scheme, solver=solver)[0].block_until_ready() first = time.perf_counter() - started started = time.perf_counter() dofsl, report = EllipticSaddleScheme.solve_pure(scheme, solver=solver) dofsl.block_until_ready() return dofsl, report, first, time.perf_counter() - started VARIANTS = { # name: (Krylov, sweep, mode, sign of the Schur inverse) "none (BiCGStab)": ("bicgstab", None, None, None), "none (MINRES)": ("minres", None, None, None), "block-diagonal + MINRES": ("minres", [U, P], "jacobi", +1.0), "block-triangular + FGMRES": ("fgmres", [P, U], "gauss_seidel", -1.0), "block LU + FGMRES": ("fgmres", [U, P, U], "gauss_seidel", -1.0), } def main(): print(f"Stokes Q2/Q1, nu = {NU}, Krylov to {TOL:.0e}\n") print("per cell: iterations | setup + first solve (compilation) | warm solve\n") header = f"{'':28s}" + "".join(f"{f'{n}x{n}':^26s}" for n in N_CELLS) print(header) rows = {} for name, (krylov, sweep, mode, sign) in VARIANTS.items(): cells = [] for n_cells in N_CELLS: scheme = stokes_scheme(make_mesh(n_cells)) if sweep is None and n_cells > 32: cells.append(f"{'--':^26s}") continue started = time.perf_counter() preconditioner = ( None if sweep is None else BlockSchurPreconditioner( sweep=sweep, inverses=inverses(scheme, n_cells, sign), mode=mode ) ) solver = LinearKrylov( tol=TOL, cg_solver=krylov, preconditioner=preconditioner, max_iter_linear=20000, ) setup = time.perf_counter() - started dofsl, report, first, warm = solve(scheme, solver) rows.setdefault(n_cells, []).append((name, errors(dofsl, scheme))) cells.append( f"{int(report.n_linear):5d} | {setup:4.1f}+{first:4.1f} s | {warm:5.2f} s " ) print(f"{name:28s}" + "".join(cells)) print( "\nDOFs:", ", ".join( f"{n}x{n}: {stokes_scheme(make_mesh(n))._initial_dofs().size}" for n in N_CELLS ), ) print("\nL2 errors (u, mean-free p), identical across preconditioners:") for n_cells, results in rows.items(): for name, (error_u, error_p) in results: print( f" {n_cells:3d}x{n_cells:<3d} {name:28s} u {error_u:.3e} p {error_p:.3e}" ) if __name__ == "__main__": main()