"""Solve a steady, periodic 2D reaction-diffusion problem with DG. -Delta u + c u = f on (0, 1)^2, periodic in both directions u_exact(x, y) = sin(2 pi x) cos(2 pi y) Periodicity is imposed by the mesh (``PeriodicEllipticDGscheme``): the faces of opposite sides are glued, so no boundary condition is needed. """ import jax.numpy as jnp import matplotlib.pyplot as plt from scimba_jax.linear_approximation.basis.analytic_bases import local_taylor_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.dg.flux import SIPGFlux from scimba_jax.linear_approximation.galerkin.dg.periodic_dg_scheme import ( PeriodicEllipticDGscheme, ) 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_dg import VariablesDG 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.diffusion_advection_reaction_weak_form import ( EllipticWeakForm, ) physical_dim = 2 quad_order = 2 out_dim = 1 mapping_id = InvertibleFunction(lambda x: x, lambda y: y) C_REACTION = 1.0 TWO_PI = 2.0 * jnp.pi def u_exact_fn(x): return (jnp.sin(TWO_PI * x[0]) * jnp.cos(TWO_PI * x[1])).reshape(out_dim) pde = EllipticWeakForm( dim=physical_dim, A=lambda x: jnp.eye(physical_dim), b=lambda x: jnp.zeros(physical_dim), c=lambda x: jnp.array(C_REACTION), f=lambda x: (2.0 * TWO_PI**2 + C_REACTION) * jnp.sin(TWO_PI * x[0]) * jnp.cos(TWO_PI * x[1]), ) def make_assembler(n_cells, order): mesh = Mesh( dim=physical_dim, n_cells=(n_cells, n_cells), ref_quad=UnitSquareTensorized(dim=physical_dim, order=quad_order), mapping=Mapping(mappings=[mapping_id]), ) taylor_basis = AnalyticBasis( nb_basis=(order + 1) ** physical_dim, out_dim=out_dim, mesh=mesh, local_basis=lambda coords, i, m, _o=order: local_taylor_basis( coords, i, m, order=_o, out_dim=out_dim ), basis_type="scalar", ) variables = VariablesDG(basis=taylor_basis, nb_variables=out_dim) model = AbstractPhysicalWeakModel.from_weak_form(pde) flux = SIPGFlux(sigma=order * (order + 1), h=1.0 / n_cells) return PeriodicEllipticDGscheme(model, variables, flux) # ── Convergence curves ────────────────────────────────────────────────────── n_cells_list = [4, 8, 16, 24] degrees = [1, 2] fig_conv, (ax_l2, ax_h1, ax_linf) = plt.subplots(1, 3, figsize=(15, 4)) last_assembler = None for order in degrees: def solver(n_cells, _o=order): assembler = make_assembler(n_cells, _o) return PeriodicEllipticDGscheme.solve(assembler) print(f"\nConvergence study p={order}:") plot_convergence( solver=solver, u_exact_fn=u_exact_fn, n_cells_list=n_cells_list, label=f"p={order}", ax_l2=ax_l2, ax_h1=ax_h1, ax_linf=ax_linf, relative=True, title_sufix=" - SIPG, periodic", ) last_assembler = solver(n_cells_list[-1]) plt.tight_layout() # ── Solution plot (finest mesh, last degree) ──────────────────────────────── n_plot = 100 grid_1d = jnp.linspace(0.0, 1.0, n_plot) xx, yy = jnp.meshgrid(grid_1d, grid_1d, indexing="ij") grid_pts = jnp.stack([xx.ravel(), yy.ravel()], axis=1) u_h = last_assembler.variables.evaluate(grid_pts)[:, 0].reshape(n_plot, n_plot) u_ex = (jnp.sin(TWO_PI * xx) * jnp.cos(TWO_PI * yy)).reshape(n_plot, n_plot) fig_sol, (ax_num, ax_ex) = plt.subplots(1, 2, figsize=(11, 4.5)) m0 = ax_num.pcolormesh(xx, yy, u_h, shading="auto", cmap="turbo") ax_num.set_title(f"DG solution (p={degrees[-1]}, n_cells={n_cells_list[-1]})") ax_num.set_xlabel("x") ax_num.set_ylabel("y") ax_num.set_aspect("equal") fig_sol.colorbar(m0, ax=ax_num, shrink=0.85) m1 = ax_ex.pcolormesh(xx, yy, u_ex, shading="auto", cmap="turbo") ax_ex.set_title("Exact solution") ax_ex.set_xlabel("x") ax_ex.set_ylabel("y") ax_ex.set_aspect("equal") fig_sol.colorbar(m1, ax=ax_ex, shrink=0.85) fig_sol.suptitle("Periodic reaction-diffusion, -Delta u + c u = f on (0,1)^2") fig_sol.tight_layout() plt.show()