r"""Solves a 2D Poisson PDE with Dirichlet boundary conditions using PINNs and FEM. .. math:: -\Delta u & = f \quad \text{in } \Omega \\ u & = 0 \quad \text{on } \partial \Omega where :math:`x = (x_1, x_2) \in \Omega` and :math:`\Omega` is a fish-shaped domain (:class:`Fish2D`: a body ellipse with a dorsal, tail and pelvic fin). Two solvers are compared on the same problem: - a PINN (simple MLP, natural-gradient optimization, boundary conditions enforced weakly); - CG-FEM on an unstructured quadrilateral mesh of the fish (Dirichlet enforced strongly on the boundary dofs), built by meshing the fish's own boundary curves (:meth:`Fish2D.full_bc_domain`) with GMSH. Both solutions are plotted, together with their pointwise difference. """ import timeit from pathlib import Path import jax import jax.numpy as jnp import matplotlib.pyplot as plt import numpy as np from scimba_jax.domains.meshless_domains.examples_domains import Fish2D from scimba_jax.linear_approximation.basis.analytic_bases import ( local_lagrange_basis, local_lagrange_basis_by_logical, ) from scimba_jax.linear_approximation.basis.dof_map import UnstructuredLagrangeDofMap 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.meshes.unstructured_mesh import UnstructuredMesh from scimba_jax.linear_approximation.quad.gauss_quad import UnitSquareTensorized from scimba_jax.linear_approximation.variables.variables_fe import VariablesFE from scimba_jax.mapping.macro_mesh import macro_mesh_from_points from scimba_jax.nonlinear_approximation.approximation_spaces.approximation_spaces import ( ApproximationSpace, ) from scimba_jax.nonlinear_approximation.integration.monte_carlo import ( DomainSampler, TensorizedSampler, ) from scimba_jax.nonlinear_approximation.networks.mlp import MLP from scimba_jax.nonlinear_approximation.numerical_solvers.projectors import Projector 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.physical_models.elliptic_pde.laplacians import LaplacianDirichletND from scimba_jax.plots.plots_galerkin import sample_solution from scimba_jax.plots.plots_nd import plot_abstract_approx_space N_COLLOC = 7500 N_BC_COLLOC = 4000 N_EPOCHS = 1500 FEM_ORDER = 3 FEM_MESH_SIZE = 0.025 N_BOUNDARY_SAMPLES = 200 _HERE = Path(__file__).resolve().parent key = jax.random.PRNGKey(0) def f_rhs(xy: jnp.ndarray) -> jnp.ndarray: return jnp.ones_like(xy[0]) * 50 dx = Fish2D(is_main_domain=True) sampler = TensorizedSampler([DomainSampler(dx)], bc=True) nn = MLP(in_size=2, out_size=1, hidden_sizes=[16] * 3, key=key) space = ApproximationSpace({"x": 2}, [(nn, "scalar", None)], model_type="x") model = LaplacianDirichletND(dx, lambda *args: f_rhs(*args), bc="weak") weights_dict = {"interior": [1.0], "boundary": [40.0]} pinn = Projector(model, space, sampler, weights=weights_dict) start = timeit.default_timer() key, pinn = pinn.project(key, space, N_EPOCHS, N_COLLOC, N_BC_COLLOC) end = timeit.default_timer() print(f"PINN training took {end - start:.1f} s") plot_abstract_approx_space( pinn.space, dx, loss=pinn.losses, draw_contours=True, n_drawn_contours=10, ) # ── FEM solution, on an unstructured mesh of the fish ──────────────────────── def fish_boundary_points(domain: Fish2D, n_per_curve: int) -> np.ndarray: """Closed polygon tracing the fish outline (tail -> nose -> tail). ``full_bc_domain`` returns the upper and lower boundary curves, both parametrized tail (``t=0``) to nose (``t=1``); tracing upper forward then lower backward gives a single ordered, closed loop, as required by :func:`macro_mesh_from_points`. """ upper, lower = domain.full_bc_domain() t_upper = jnp.linspace(0.0, 1.0, n_per_curve)[:, None] t_lower = jnp.linspace(1.0, 0.0, n_per_curve)[:, None] upper_pts = np.asarray(upper.surface_o_mapping(t_upper)) lower_pts = np.asarray(lower.surface_o_mapping(t_lower)) # drop lower's endpoints: they duplicate upper's nose (t=1) and the # implicit closing edge back to upper's tail (t=0). return np.concatenate([upper_pts, lower_pts[1:-1]], axis=0) def cached_fish_mesh( domain: Fish2D, order: int, mesh_size: float, n_per_curve: int ) -> tuple[np.ndarray, np.ndarray]: """Mesh the fish once and cache it: GMSH does not repeat itself across calls.""" cache = _HERE / f"mesh_fish_p{order}_h{mesh_size}.npz" if cache.exists(): stored = np.load(cache) return stored["nodes"], stored["cells"] boundary = fish_boundary_points(domain, n_per_curve) macro = macro_mesh_from_points( boundary, order=order, mesh_size=mesh_size, smooth=True ) np.savez_compressed(cache, nodes=macro.nodes, cells=macro.cells) print(f" meshed the fish and saved {cache.name}") return macro.nodes, macro.cells def solve_laplacian_fe(mesh: UnstructuredMesh, order: int) -> EllipticFEscheme: """Solve ``-Delta u = 50`` with ``u = 0`` on the whole boundary, by CG-FEM.""" basis = AnalyticBasis( nb_basis=(order + 1) ** 2, out_dim=1, mesh=mesh, local_basis_by_logical=lambda y, i, m: local_lagrange_basis_by_logical( y, i, m, order=order, out_dim=1 ), local_basis=lambda y, i, m: local_lagrange_basis( y, i, m, order=order, out_dim=1 ), basis_type="scalar", ) fe_model = AbstractPhysicalWeakModel.from_weak_form( LaplacianWeakForm(dim=2, f=f_rhs), dirichlet=lambda x: jnp.zeros(1), ) variables = VariablesFE( basis=basis, nb_variables=1, dof_map=UnstructuredLagrangeDofMap ) return EllipticFEscheme.solve(EllipticFEscheme(fe_model, variables), max_iter=1) start = timeit.default_timer() nodes, cells = cached_fish_mesh(dx, FEM_ORDER, FEM_MESH_SIZE, N_BOUNDARY_SAMPLES) fem_mesh = UnstructuredMesh( nodes=nodes, cells=cells, ref_quad=UnitSquareTensorized(dim=2, order=2 * FEM_ORDER + 2), order=FEM_ORDER, ) fem = solve_laplacian_fe(fem_mesh, FEM_ORDER) jax.block_until_ready(fem.variables.dofsl) end = timeit.default_timer() print(f"FEM solve took {end - start:.1f} s on {fem_mesh.n_cells_total} cells") # ── Compare the two solutions, sampled at the same points ──────────────────── points, triangles, u_fem = sample_solution(fem, n_side=8) u_pinn = np.asarray(pinn.evaluate(jnp.asarray(points))[:, 0]) error = np.log10(np.abs(u_pinn - u_fem)) figure, axes = plt.subplots(1, 3, figsize=(16, 5)) for ax, values, title, cmap in zip( axes, [u_pinn, u_fem, error], ["PINN (weak BC)", "FEM (strong BC)", "log(|PINN - FEM|)"], ["turbo", "turbo", "magma"], ): drawing = ax.tricontourf( points[:, 0], points[:, 1], triangles, values, levels=40, cmap=cmap ) figure.colorbar(drawing, ax=ax, fraction=0.046) ax.set_title(title) ax.set_aspect("equal") ax.set_xlabel("x") ax.set_ylabel("y") figure.suptitle("-Δu = 50 on the fish domain: PINN vs FEM") figure.tight_layout() print(f"max|PINN - FEM| = {error.max():.4f}, mean|PINN - FEM| = {error.mean():.4f}") plt.show()