"""Solve the 2D Laplacian FEM scheme on a single mesh — matrix-free solver. Exact solution: u(x, y) = sin(π x) sin(π y) on [0, 1]² PDE: -Δu = 2 π² sin(π x) sin(π y), u = 0 on ∂Ω The linear system J δu = -F is solved matrix-free: the Jacobian is never assembled, only its action J·v (a JVP) is provided to a Krylov solver (CG, the operator being SPD for pure diffusion). """ import time import jax import jax.numpy as jnp from scimba_jax.linear_approximation.basis.analytic_bases import local_lagrange_basis from scimba_jax.linear_approximation.basis.general_bases import AnalyticBasis from scimba_jax.linear_approximation.error_analysis import ( h1_seminorm_error, l2_error, linf_error, ) from scimba_jax.linear_approximation.galerkin.fem.elliptic_fe_scheme import ( EllipticFEscheme, ) from scimba_jax.linear_approximation.meshes.mesh import Mesh from scimba_jax.linear_approximation.quad.gauss_quad import UnitSquareTensorized 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.laplacian_weak_form import ( LaplacianWeakForm, ) from scimba_jax.plots.plots_galerkin import plot_solution_2d # ── Problem setup ───────────────────────────────────────────────────────────── physical_dim = 2 quad_order = 3 out_dim = 1 basis_order = 1 # Q1 elements n_cells = 64 # single mesh mapping_id = InvertibleFunction(lambda x: x, lambda y: y) pde = LaplacianWeakForm( dim=physical_dim, f=lambda x: 2 * jnp.pi**2 * jnp.sin(jnp.pi * x[0]) * jnp.sin(jnp.pi * x[1]), ) def u_exact_fn(x): return (jnp.sin(jnp.pi * x[0]) * jnp.sin(jnp.pi * x[1])).reshape(out_dim) def dirichlet_bc(__x): return jnp.zeros(out_dim) def make_assembler(): m = Mesh( dim=physical_dim, n_cells=(n_cells, n_cells), ref_quad=UnitSquareTensorized(dim=physical_dim, order=quad_order), mapping=Mapping(mappings=[mapping_id]), ) lagrange_basis = AnalyticBasis( nb_basis=(basis_order + 1) ** physical_dim, out_dim=out_dim, mesh=m, local_basis=lambda coords, i, mesh: local_lagrange_basis( coords, i, mesh, order=basis_order, out_dim=out_dim ), basis_type="scalar", ) variables = VariablesFE(basis=lagrange_basis, nb_variables=out_dim) return EllipticFEscheme( AbstractPhysicalWeakModel.from_weak_form(pde, dirichlet=dirichlet_bc), variables ) def report(tag, assembler): l2 = float(l2_error(assembler, u_exact_fn, relative=True)) h1 = float(h1_seminorm_error(assembler, u_exact_fn, relative=True)) linf = float(linf_error(assembler, u_exact_fn, relative=True)) print(f" [{tag}] L2={l2:.3e} H1={h1:.3e} Linf={linf:.3e}") return l2, h1, linf # ── Matrix-free solve ───────────────────────────────────────────────────────── print(f"2D Laplacian CG-FEM — {n_cells}×{n_cells} mesh, Q{basis_order}") ndof = make_assembler().variables.dofsl.size print(f" ndof = {ndof}") print("\nMatrix-free solve (CG):") assembler_mf = make_assembler() # Warmup: triggers JIT compilation. assembler_mf = EllipticFEscheme.solve( assembler_mf, matrix_free=True, max_iter=1, tol=1e-10 ) jax.block_until_ready(assembler_mf.variables.dofsl) # Reset and solve again for a clean execution timing. assembler_mf.variables.dofsl = jnp.zeros_like(assembler_mf.variables.dofsl) t0 = time.perf_counter() assembler_mf = EllipticFEscheme.solve( assembler_mf, matrix_free=True, max_iter=1, tol=1e-10 ) jax.block_until_ready(assembler_mf.variables.dofsl) print(f" exec = {time.perf_counter() - t0:.3f}s") report("matrix-free", assembler_mf) # ── Reference: assembled-Jacobian solve ─────────────────────────────────────── print("\nAssembled-Jacobian solve (reference):") assembler_ref = make_assembler() assembler_ref = EllipticFEscheme.solve(assembler_ref, max_iter=1) report("assembled", assembler_ref) diff = float( jnp.linalg.norm(assembler_mf.variables.dofsl - assembler_ref.variables.dofsl) / jnp.linalg.norm(assembler_ref.variables.dofsl) ) print(f"\n ‖u_mf - u_ref‖ / ‖u_ref‖ = {diff:.3e}") # ── Solution plot ────────────────────────────────────────────────────────────── print("\nPlotting matrix-free solution...") plot_solution_2d(assembler_mf, u_exact_fn)