r"""Gray-Scott reaction-diffusion system in 2D with TimeDiscreteFEscheme. .. math:: \partial_t u &= \varepsilon_1 \Delta u + b_1(1 - u) - c_1 u v^2, \\ \partial_t v &= \varepsilon_2 \Delta v - b_2 v + c_2 u v^2, on :math:`(x, y) \in (-1, 1)^2`, with initial conditions .. math:: u_0(x, y) &= 1 - \exp\!\bigl(-10\,((x+0.05)^2 + (y+0.02)^2)\bigr), \\ v_0(x, y) &= \exp\!\bigl(-10\,((x-0.05)^2 + (y-0.02)^2)\bigr). Same test case (domain, ICs, coefficients) as ``examples_jax/pinns/time_dependent_pdes/parabolic_systems/gray_scott_2d.py`` and as the DG counterpart of this file (``examples_jax/dg/solve/time_dependent/2d_gray_scott_reaction_diffusion.py``) -- kept unchanged so all three are comparable -- but solved by classical CG-FEM time-marching. Validated against an independent RK4 finite-difference solution (no analytical solution exists for this system), the same way the PINN example validates itself. Per-task, this is a single mesh/degree/time-integrator check, not a convergence sweep -- see the benchopt benchmark ``benchmarks/benchmarks_jax/heat_time_convergence`` for that. **Boundary condition: natural (zero-flux Neumann), not periodic.** The DG counterpart uses a genuinely periodic domain (torus), via :class:`~scimba_jax.linear_approximation.galerkin.dg.periodic_dg_scheme.PeriodicEllipticDGscheme` -- but no periodic scheme exists for FEM (:mod:`scimba_jax.linear_approximation.galerkin.fem` has none). ``dirichlet= None`` here registers no boundary condition at all, which is CG-FEM's *natural* BC: dropping the boundary term integration by parts leaves behind is exactly imposing zero diffusive flux through :math:`\partial\Omega`. This is close enough to the DG file's periodic reference over this short a time horizon -- both Gaussian blobs sit well inside :math:`(-1, 1)^2` and never reach the boundary -- but is a genuine (small) discrepancy with the FD reference below, which stays periodic to match the PINN test case; it shows up as a slightly larger error near the domain edges in the error plot. **The coupled system needs one fix to the framework, applied upstream.** ``u``/``v`` live on one ``out_dim=2`` ("vec") FE space -- the same pattern the coupled stationary test in ``test_solve_laplacian_fem.py`` uses, and the *nonlinear* reaction term (depending on the trial functions ``u1, u2`` themselves) goes into ``bilinear_form``, not ``linear_form``. Both this file and the DG counterpart now share the library's :class:`~scimba_jax.physical_models.classical_weakform.gray_scott_weak_form. GrayScottWeakForm` -- the weak form itself is scheme-agnostic (``FEM``/``DG`` both go through ``AbstractWeakForm``/``get_fields()``), only the assembler (``TimeDiscreteFEscheme`` here, ``TimeDiscreteDGscheme`` + a numerical flux there) differs. Working this example out surfaced a real bug in :mod:`scimba_jax.linear_approximation.galerkin.fem.time_dependent_fe_scheme`: its internal ``_mass_form`` helper summed a vector system's ``(u_i, v_i)`` pairs into *one scalar* instead of keeping one entry per component, which silently cross-wires every component's mass into every other's residual row the moment ``out_dim > 1``, as soon as the result is added to the (correctly vector-valued) spatial term. Every existing test happened to only exercise scalar (``out_dim=1``) systems through this class, so nothing caught it. Fixed at the source (now returns a proper per-component ``ParamVecFunction``, mirroring every other coupled-system weak form in this codebase) rather than routed around here. The reaction coefficients (``c1 = c2 = 1000``) make this a genuinely stiff system -- explicit RK schemes would need a prohibitively small ``dt`` for stability having nothing to do with accuracy. An L-stable, 2-stage SDIRK (Pareschi-Russo, one of the implicit tableaux of the ``heat_time_convergence`` benchmark) is used here, fully implicit (diffusion *and* reaction Newton-solved together each stage -- this codebase's time-dependent Galerkin schemes reject IMEX tableaux outright). """ 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.galerkin.fem.time_discrete_fe_scheme import ( TimeDiscreteFEscheme, ) from scimba_jax.linear_approximation.meshes.cartesian_mesh import cartesian_mesh from scimba_jax.linear_approximation.variables.variables_fe import VariablesFE from scimba_jax.physical_models.classical_weakform.gray_scott_weak_form import ( GrayScottWeakForm, gray_scott_fd_grid, gray_scott_fd_reference, ) from scimba_jax.time_discrete.butcher_tableau import build_pareschi_russo_tableau jax.config.update("jax_enable_x64", True) DIM = 2 # ── Problem parameters (identical to the PINN reference test case) ─────────── EPS1 = 0.2 # D_u EPS2 = 0.1 # D_v B1 = 40.0 # feed rate F B2 = 100.0 # F + k C1 = 1000.0 # reaction coefficient for u C2 = 1000.0 # reaction coefficient for v A_LO, A_HI = -1.0, 1.0 # domain (-1, 1)^2 BOX_BOUNDS = ((A_LO, A_HI), (A_LO, A_HI)) def u0v0(xy): """Initial condition, pointwise: (x, y) -> (u0, v0).""" x, y = xy[0], xy[1] u0 = 1.0 - jnp.exp(-10.0 * ((x + 0.05) ** 2 + (y + 0.02) ** 2)) v0 = jnp.exp(-10.0 * ((x - 0.05) ** 2 + (y - 0.02) ** 2)) return jnp.stack([u0, v0]) # ── FE space and solve ──────────────────────────────────────────────────────── def make_variables(n_cells: int, order: int, quad_order: int) -> VariablesFE: mesh = cartesian_mesh( n_cells=(n_cells, n_cells), quad_order=quad_order, bounds=BOX_BOUNDS ) basis = AnalyticBasis( nb_basis=(order + 1) ** DIM, out_dim=2, mesh=mesh, local_basis=lambda y, i, m, k=order: local_lagrange_basis( y, i, m, order=k, out_dim=2 ), basis_type="vec", ) return VariablesFE(basis=basis, nb_variables=2) def solve(n_cells: int, order: int, dt: float, nt: int): quad_order = order + 3 variables = make_variables(n_cells, order, quad_order) weak_form = GrayScottWeakForm( dim=DIM, D_u=EPS1, D_v=EPS2, b1=B1, b2=B2, c1=C1, c2=C2 ) scheme = TimeDiscreteFEscheme( spatial_weak_form_factory=weak_form, variables=variables, butcher_tableau=build_pareschi_russo_tableau(), dt=dt, dirichlet=None, cg_solver="bicgstab", max_iter=30, tol=1e-10, ) dofsl_init = scheme.initialize(u0v0) dofsl_final, history = scheme.solve(dofsl_init, t0=0.0, nt=nt) return variables, dofsl_final, history # ── Reference: RK4 finite-difference solve on a fine periodic grid ─────────── # Same reference the DG counterpart and the PINN example use -- see # gray_scott_fd_reference in the library for why it lives there. Kept # periodic (rather than Neumann, which would match the FEM BC exactly) so # this reference is the same for both example files -- see the module # docstring for why the mismatch does not matter here. N_FD = 160 _x_fd, _xy_fd = gray_scott_fd_grid((A_LO, A_HI), N_FD) def fd_reference(t_final: float, dt_fd: float = 2e-5): return gray_scott_fd_reference( u0v0, (A_LO, A_HI), EPS1, EPS2, B1, B2, C1, C2, t_final, N_FD, dt_fd ) if __name__ == "__main__": N_CELLS = 24 ORDER = 2 DT = 1e-2 NT = 40 T_FINAL = DT * NT print( f"Solving Gray-Scott with FEM: n_cells={N_CELLS}, order={ORDER}, " f"dt={DT:.1e}, nt={NT} (T_final={T_FINAL:.4f}) -- Pareschi-Russo, Neumann" ) t0 = time.time() variables, dofsl_final, _history = solve(N_CELLS, ORDER, DT, NT) print(f"FEM solve done in {time.time() - t0:.1f}s") print(f"Solving independent RK4/FD reference on a {N_FD}x{N_FD} periodic grid...") t0 = time.time() u_ref, v_ref = fd_reference(T_FINAL) print(f"FD reference done in {time.time() - t0:.1f}s") variables.dofsl = dofsl_final uv_fem = variables.evaluate(_xy_fd) u_fem = np.array(uv_fem[:, 0]).reshape(N_FD, N_FD) v_fem = np.array(uv_fem[:, 1]).reshape(N_FD, N_FD) err_u = np.abs(u_fem - u_ref) err_v = np.abs(v_fem - v_ref) rel_l2_u = np.linalg.norm(err_u) / np.linalg.norm(u_ref) rel_l2_v = np.linalg.norm(err_v) / np.linalg.norm(v_ref) print( f"u: rel. L2 error vs. FD reference = {rel_l2_u:.3e}, max abs = {err_u.max():.3e}" ) print( f"v: rel. L2 error vs. FD reference = {rel_l2_v:.3e}, max abs = {err_v.max():.3e}" ) # ── Plot: FEM vs. FD reference vs. |difference| ─────────────────────────── X, Y = np.array(_x_fd), np.array(_x_fd) fig, axes = plt.subplots(2, 3, figsize=(15, 9)) panels = [ (u_fem, "$u$ FEM"), (u_ref, "$u$ FD reference (periodic)"), (err_u, r"$|u_\mathrm{FEM} - u_\mathrm{FD}|$"), (v_fem, "$v$ FEM"), (v_ref, "$v$ FD reference (periodic)"), (err_v, r"$|v_\mathrm{FEM} - v_\mathrm{FD}|$"), ] for ax, (grid, title) in zip(axes.flat, panels): im = ax.pcolormesh(X, Y, grid, cmap="turbo", shading="auto") fig.colorbar(im, ax=ax, pad=0.02) ax.set_title(title) ax.set_xlabel("x") ax.set_ylabel("y") ax.set_aspect("equal") fig.suptitle( rf"Gray-Scott 2D, FEM (Pareschi-Russo, Q{ORDER}, {N_CELLS}x{N_CELLS} cells, " rf"Neumann) vs. FD reference (periodic) at $t={T_FINAL:.4f}$" ) fig.tight_layout() plt.show()