"""Solve the Laplacian FEM scheme via Newton-Raphson.""" import jax.numpy as jnp import matplotlib.pyplot as plt 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 plot_convergence 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, ) physical_dim = 1 quad_order = 3 out_dim = 1 mapping_id = InvertibleFunction(lambda x: x, lambda y: y) pde = LaplacianWeakForm( dim=physical_dim, f=lambda x: jnp.sin(jnp.pi * x[0]), ) def u_exact_fn(x): return jnp.sin(jnp.pi * x[0]) / jnp.pi**2 # ── Convergence curves ──────────────────────────────────────────────────────── n_cells_list = [5, 10, 20, 40, 80] fig_conv, axes = plt.subplots(1, 3, figsize=(15, 4), squeeze=False) ax_l2, ax_h1, ax_linf = axes[0] for deg in [1, 2]: basis_order = deg def make_solver(b=basis_order): def solver(n_cells): m = Mesh( dim=physical_dim, n_cells=[n_cells], ref_quad=UnitSquareTensorized(dim=physical_dim, order=quad_order), mapping=Mapping(mappings=[mapping_id]), ) lagrange_basis = AnalyticBasis( nb_basis=b + 1, out_dim=out_dim, mesh=m, local_basis=lambda coords, i, mesh, _b=b: local_lagrange_basis( coords, i, mesh, order=_b, out_dim=out_dim ), basis_type="scalar", ) variables = VariablesFE(basis=lagrange_basis, nb_variables=out_dim) def dirichlet_bc(__x): return jnp.zeros(out_dim) assembler = EllipticFEscheme( AbstractPhysicalWeakModel.from_weak_form(pde, dirichlet=dirichlet_bc), variables, ) return EllipticFEscheme.solve(assembler) return solver print(f"\nConvergence study FEM p={deg}:") plot_convergence( solver=make_solver(), u_exact_fn=u_exact_fn, n_cells_list=n_cells_list, label=f"p={deg}", ax_l2=ax_l2, ax_h1=ax_h1, ax_linf=ax_linf, relative=True, ) plt.tight_layout() plt.show()