"""phi-FEM Dirichlet solver on a circular domain — convergence study. Exact solution: u(x,y) = phi(x,y) * exp(x) * sin(2*pi*y) Level-set: phi(x,y) = (x-0.5)^2 + (y-0.5)^2 - R^2 (R = sqrt(2)/4) The implicit Dirichlet BC u = 0 on Gamma is encoded automatically via u_h = phi_h * w_h. """ import time import jax import jax.numpy as jnp import matplotlib.pyplot as plt import numpy as np 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.meshes.mesh import Mesh from scimba_jax.linear_approximation.phi_fem.poisson_dirichlet_phifem_scheme import ( PoissonDirichletDirectPhiFEscheme, ) 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.classical_weakform.laplacian_weak_form import ( LaplacianWeakForm, ) # ── Problem setup ────────────────────────────────────────────────────────────── R = 2 ** (1 / 2) / 4 cx, cy = 0.5, 0.5 def phi(x): return (x[0] - cx) ** 2 + (x[1] - cy) ** 2 - R**2 def u_exact(x): return phi(x) * jnp.exp(x[0]) * jnp.sin(2 * jnp.pi * x[1]) def f_source(x): H = jax.hessian(u_exact)(x) return -jnp.trace(H) # ── Parameters ───────────────────────────────────────────────────────────────── dim = 2 quad_order = 3 out_dim = 1 mapping_id = InvertibleFunction(lambda x: x, lambda y: y) # ── Convergence study ────────────────────────────────────────────────────────── n_cells_list = [10] basis_order = 1 # fig_conv, axes = plt.subplots(1, 3, figsize=(15, 4)) # ax_l2, ax_h1, ax_linf = axes def make_solver(b=basis_order): def solver(n_cells): mesh = Mesh( dim=dim, n_cells=[n_cells, n_cells], ref_quad=UnitSquareTensorized(dim=dim, order=quad_order), mapping=Mapping(mappings=[mapping_id]), ) lagrange_basis = AnalyticBasis( nb_basis=(b + 1) ** dim, out_dim=out_dim, mesh=mesh, 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) pde = LaplacianWeakForm(dim=dim, f=f_source) scheme = PoissonDirichletDirectPhiFEscheme(pde, variables, phi, sigma_D=1.0) return PoissonDirichletDirectPhiFEscheme.solve(scheme) return solver # print(f"\nConvergence study phi-FEM p={basis_order}:") # plot_convergence( # solver=make_solver(), # u_exact_fn=u_exact, # n_cells_list=n_cells_list, # label=f"p={basis_order}", # ax_l2=ax_l2, # ax_h1=ax_h1, # ax_linf=ax_linf, # relative=True, # ) # plt.suptitle("phi-FEM convergence — circular domain") # plt.tight_layout() # plt.show() # ── Solution plot (last mesh) ────────────────────────────────────────────────── n_cells = n_cells_list[-1] t0 = time.perf_counter() scheme = make_solver(b=basis_order)(n_cells) print(f"[example] make_solver (n={n_cells}): {time.perf_counter() - t0:.3f}s") n_eval = 50 xs = np.linspace(0.0, 1.0, n_eval) ys = np.linspace(0.0, 1.0, n_eval) XX, YY = np.meshgrid(xs, ys) pts = jnp.array(np.stack([XX.ravel(), YY.ravel()], axis=-1)) # Mask: point belongs to an active cell (interior or cut) # Cell ordering: cell (i, j) -> index i * n_cells + j # i = x-direction index, j = y-direction index ci = np.minimum((XX * n_cells).astype(int), n_cells - 1) cj = np.minimum((YY * n_cells).astype(int), n_cells - 1) cell_idx_grid = ci * n_cells + cj active_mask = np.isin(cell_idx_grid, scheme.active_cells).ravel() pts_active = pts[active_mask] u_h_vals = np.array(scheme.evaluate_u(pts_active)) u_ex_vals = np.array(jnp.vectorize(u_exact, signature="(n)->()")(pts_active)) # linf_err = float(scheme.linf_error(u_exact)) l2_err = float(scheme.l2_error(u_exact, relative=True)) fig, axes = plt.subplots(1, 3, figsize=(15, 4)) U_h = np.full(XX.shape, np.nan) U_ex = np.full(XX.shape, np.nan) U_err = np.full(XX.shape, np.nan) U_h.ravel()[active_mask] = u_h_vals U_ex.ravel()[active_mask] = u_ex_vals U_err.ravel()[active_mask] = np.abs(u_h_vals - u_ex_vals) im0 = axes[0].contourf(XX, YY, U_h, levels=20, cmap="RdBu_r") axes[0].set_title("$u_h = \\varphi_h w_h$ (phi-FEM)") plt.colorbar(im0, ax=axes[0]) im1 = axes[1].contourf(XX, YY, U_ex, levels=20, cmap="RdBu_r") axes[1].set_title("$u_{exact}$") plt.colorbar(im1, ax=axes[1]) im2 = axes[2].contourf(XX, YY, U_err, levels=20, cmap="Reds") axes[2].set_title(f"|u_h - u_exact| (L² = {l2_err:.2e})") plt.colorbar(im2, ax=axes[2]) for ax in axes: theta = np.linspace(0, 2 * np.pi, 300) ax.plot(cx + R * np.cos(theta), cy + R * np.sin(theta), "k--", lw=1) ax.set_aspect("equal") plt.suptitle("phi-FEM P1 — 20×20 mesh") plt.tight_layout() plt.show()