"""Cell-centred finite volume for ``-u'' = 1`` on ``[0, 1]``. The script solves the same problem twice: with the dense autodiff Jacobian and with a matrix-free Krylov solve preconditioned by the FV point Jacobi. ``VariablesFV(..., projection="center")`` would initialise a field by its cell-centre values; ``"average"`` (the default) uses cell quadrature. It also displays the cell solution and the observed spatial convergence curve. """ import jax import jax.numpy as jnp import matplotlib.pyplot as plt from scimba_jax.linear_approximation.finite_volume import ( EllipticFVScheme, FVJacobiPreconditioner, TPFADiffusionFlux, ) from scimba_jax.linear_approximation.meshes.cartesian_mesh import cartesian_mesh from scimba_jax.linear_approximation.solvers import LinearKrylov, LinearSolve from scimba_jax.linear_approximation.variables.variables_fv import VariablesFV from scimba_jax.nonlinear_approximation.model_class.funcparam_matrix import ( ParamMatrixFunction, ) from scimba_jax.nonlinear_approximation.model_class.funcparam_scalar import ( ParamScalarFunction, ) from scimba_jax.physical_models.abstract_conservative_pde import ( AbstractConservativePDE, ) from scimba_jax.physical_models.abstract_physical_conservative_model import ( AbstractPhysicalConservativeModel, ) class Poisson1D(AbstractConservativePDE): """``div(F) - div(A grad(u)) + R = f`` with ``F=R=0, A=1, f=1``.""" def __init__(self): super().__init__(dim=1) def construct_F(self, u): # noqa: N802 return None def construct_A(self, u): # noqa: N802 return ParamMatrixFunction(u.dims, lambda _sp, _x: jnp.eye(1), f_type=u.f_type) def construct_R(self, u): # noqa: N802 return None def construct_f(self): return ParamScalarFunction( {"x": 1, "mu": 0}, lambda _sp, _x: jnp.array(1.0), f_type="x" ) def exact_solution(x): """Exact solution, evaluated at a physical point or batch of points.""" return 0.5 * x[..., :1] * (1.0 - x[..., :1]) def make_scheme(n_cells=32): mesh = cartesian_mesh([n_cells], quad_order=3) model = AbstractPhysicalConservativeModel.from_pde( Poisson1D(), dirichlet=lambda _x: jnp.zeros(1) ) return EllipticFVScheme( model, VariablesFV(mesh, projection="average"), TPFADiffusionFlux(), assemble_first_order=False, assemble_reaction=False, ) def solve(matrix_free, n_cells=32): """Solve one formulation and return the scheme and its solve report.""" scheme = make_scheme(n_cells) solver = ( LinearKrylov( tol=1e-12, preconditioner=FVJacobiPreconditioner(), max_iter_linear=128, ) if matrix_free else LinearSolve(tol=1e-12) ) return EllipticFVScheme.solve(scheme, solver=solver, return_report=True) def cell_centres(scheme): """Physical centres of all FV cells.""" return jax.vmap(scheme.variables.mesh.cell_centroid)( jnp.arange(scheme.variables.mesh.n_cells_total) ) def max_error(scheme): """Maximum error against the manufactured solution at cell centres.""" centers = cell_centres(scheme) return float(jnp.max(jnp.abs(scheme.variables.dofsl - exact_solution(centers)))) def convergence_study(reference_scheme, n_cells=(8, 16, 32)): """Observed cell-centre convergence of the matrix-free discretisation.""" errors = [] for n_cells_one_dim in n_cells: if n_cells_one_dim == reference_scheme.variables.mesh.n_cells[0]: scheme = reference_scheme else: scheme, _ = solve(matrix_free=True, n_cells=n_cells_one_dim) errors.append(max_error(scheme)) return jnp.asarray(1.0 / jnp.asarray(n_cells)), jnp.asarray(errors) def print_summary(dense, dense_report, matrix_free, mf_report, errors): """Print a compact, human-readable solver summary.""" print("\nFinite volume 1D — -u'' = 1, u(0) = u(1) = 0") print("=" * 58) print(f"mailles : {matrix_free.variables.mesh.n_cells_total}") print(f"solveur dense : {int(dense_report.n_iter)} résolution") print(f"résidu dense : {float(dense_report.residual):.3e}") print(f"solveur matrix-free : {int(mf_report.n_iter)} résolution") print(f"itérations Krylov : {int(mf_report.n_linear)}") print(f"résidu matrix-free : {float(mf_report.residual):.3e}") print(f"erreur max. au centre : {max_error(matrix_free):.3e}") print( "écart dense / MF : " f"{float(jnp.max(jnp.abs(dense.variables.dofsl - matrix_free.variables.dofsl))):.3e}" ) rates = jnp.log(errors[:-1] / errors[1:]) / jnp.log(2.0) print(f"ordre observé : {float(jnp.mean(rates)):.2f}") def plot_results(dense, matrix_free, mesh_sizes, errors): """Show the numerical solution and the mesh-convergence curve.""" centers = cell_centres(matrix_free) exact = exact_solution(centers) figure, (solution_ax, convergence_ax) = plt.subplots(1, 2, figsize=(12, 4.5)) solution_ax.plot(centers[:, 0], exact[:, 0], "k-", lw=2.5, label="exacte") solution_ax.plot( centers[:, 0], dense.variables.dofsl[:, 0], "o", ms=4, label="dense" ) solution_ax.plot( centers[:, 0], matrix_free.variables.dofsl[:, 0], "x", ms=5, label="matrix-free + Jacobi", ) solution_ax.set( xlabel="x", ylabel="u", title="Solution aux centres des mailles", ) solution_ax.grid(alpha=0.25) solution_ax.legend() convergence_ax.loglog(mesh_sizes, errors, "o-", lw=2, label="erreur max.") reference = errors[-1] * (mesh_sizes / mesh_sizes[-1]) ** 2 convergence_ax.loglog(mesh_sizes, reference, "k--", label=r"référence $h^2$") convergence_ax.set( xlabel="taille de maille h", ylabel="erreur max. aux centres", title="Convergence en maillage (matrix-free)", ) convergence_ax.grid(which="both", alpha=0.25) convergence_ax.legend() figure.tight_layout() return figure if __name__ == "__main__": dense, dense_report = solve(matrix_free=False) matrix_free, mf_report = solve(matrix_free=True) mesh_sizes, errors = convergence_study(matrix_free) print_summary(dense, dense_report, matrix_free, mf_report, errors) plot_results(dense, matrix_free, mesh_sizes, errors) plt.show()