"""Cell-centred finite volume for a 2D Dirichlet Poisson problem. It runs a dense linear solve and the matrix-free Krylov alternative with the finite-volume point-Jacobi preconditioner. The manufactured solution is ``x(1-x)y(1-y)`` on the unit square. The script also displays the numerical solution, the error field and a mesh-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 Poisson2D(AbstractConservativePDE): """``-Delta(u) = 2 (x(1-x) + y(1-y))``.""" def __init__(self): super().__init__(dim=2) 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(2), f_type=u.f_type) def construct_R(self, u): # noqa: N802 return None def construct_f(self): return ParamScalarFunction( {"x": 2, "mu": 0}, lambda _sp, x: 2.0 * (x[0] * (1.0 - x[0]) + x[1] * (1.0 - x[1])), f_type="x", ) def exact_solution(x): return x[..., 0] * (1.0 - x[..., 0]) * x[..., 1] * (1.0 - x[..., 1]) def make_scheme(n_cells=24): mesh = cartesian_mesh([n_cells, n_cells], quad_order=3) model = AbstractPhysicalConservativeModel.from_pde( Poisson2D(), 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=24): """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=256, ) 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 rms_error(scheme): """RMS error against the manufactured solution at the cell centres.""" centers = cell_centres(scheme) exact = exact_solution(centers)[:, None] return float(jnp.sqrt(jnp.mean((scheme.variables.dofsl - exact) ** 2))) def convergence_study(reference_scheme, n_cells=(8, 16, 24)): """Observed spatial convergence of the matrix-free FV solve.""" 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(rms_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 2D — -Δu = 2[x(1-x) + y(1-y)]") print("=" * 58) print(f"maillage : {matrix_free.variables.mesh.n_cells.tolist()}") 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 RMS au centre : {rms_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(jnp.asarray([2.0, 1.5])) print(f"ordre observé : {float(jnp.mean(rates)):.2f}") def _cell_scatter( axis, centers, values, title, n_cells, vmin=None, vmax=None, cmap="viridis" ): """Draw piecewise-constant cell data at a size suited to the mesh.""" image = axis.scatter( centers[:, 0], centers[:, 1], c=values, marker="s", s=18000.0 / n_cells, linewidths=0, vmin=vmin, vmax=vmax, cmap=cmap, ) axis.set(xlabel="x", ylabel="y", title=title, aspect="equal") return image def plot_results(dense, matrix_free, mesh_sizes, errors): """Show exact, dense, matrix-free, errors and mesh convergence.""" centers = cell_centres(matrix_free) exact = exact_solution(centers) dense_numerical = dense.variables.dofsl[:, 0] numerical = matrix_free.variables.dofsl[:, 0] dense_error = jnp.abs(dense_numerical - exact) matrix_free_error = jnp.abs(numerical - exact) n_cells = matrix_free.variables.mesh.n_cells[0] figure, axes = plt.subplots(2, 3, figsize=(14, 8.5), constrained_layout=True) lo, hi = float(jnp.min(exact)), float(jnp.max(exact)) image = _cell_scatter( axes[0, 0], centers, exact, "Solution exacte", n_cells, lo, hi ) _cell_scatter( axes[0, 1], centers, dense_numerical, "Solution VF dense", n_cells, lo, hi ) _cell_scatter( axes[0, 2], centers, numerical, "Solution VF matrix-free", n_cells, lo, hi ) figure.colorbar(image, ax=axes[0, :], shrink=0.82, label="u") max_error = float(jnp.maximum(jnp.max(dense_error), jnp.max(matrix_free_error))) error_image = _cell_scatter( axes[1, 0], centers, dense_error, "Erreur absolue dense", n_cells, vmin=0.0, vmax=max_error, cmap="magma", ) _cell_scatter( axes[1, 1], centers, matrix_free_error, "Erreur absolue matrix-free", n_cells, vmin=0.0, vmax=max_error, cmap="magma", ) figure.colorbar(error_image, ax=axes[1, :2], shrink=0.82, label="|u - u_h|") convergence_ax = axes[1, 2] convergence_ax.loglog(mesh_sizes, errors, "o-", lw=2, label="erreur RMS") 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 RMS aux centres", title="Convergence en maillage (matrix-free)", ) convergence_ax.grid(which="both", alpha=0.25) convergence_ax.legend() 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()