r"""1D nonlinear reaction-diffusion, with TimeDiscreteCollocationScheme. .. math:: \partial_t u - \partial_{xx} u + u^2 = f(x, t) \quad \text{on } (0, 1), \qquad u = 0 \text{ on the boundary}, the simplest genuinely nonlinear extension of this directory's ``1d_heat_equation.py`` (same Laplacian, one quadratic reaction term added) -- a scalar, 1D, single-term nonlinearity, deliberately not a stiff or multi-component problem like ``examples_jax/dg/solve/time_dependent/ 2d_gray_scott_reaction_diffusion.py``. **Manufactured solution**, same idea as this directory's own ``2d_advection_reaction_diffusion_isotropic_diffusion.py`` (there, a linear ADR operator; here, ``u^2`` too): pick .. math:: u_\mathrm{exact}(x, t) = e^{-t} \sin(\pi x), which is 0 on the whole boundary for every :math:`t` (so the boundary condition stays homogeneous) and matches the initial condition :math:`u(x, 0) = \sin(\pi x)` at :math:`t = 0`, and solve for the :math:`f` that makes the PDE exact for that :math:`u`: .. math:: \partial_t u &= -u, \\ -\partial_{xx} u &= \pi^2 u, \\ u^2 &= e^{-2t} \sin^2(\pi x), so :math:`f = (\pi^2 - 1)\,u + u^2`. **Not this codebase's own ``AllenCahnResidual``/``GrayScottResidual``** (:mod:`scimba_jax.physical_models.temporal_pde`): those are written for continuous-time-and-space PINN collocation -- they call :meth:`~scimba_jax.nonlinear_approximation.model_class.funcparam_matrix. ParamScalarFunction.d_t` *inside* ``construct_residual`` and take a ``time_domain`` (``model_type="t_x"``) -- not the first-order-in-time, Butcher-tableau time-marching :class:`~scimba_jax.linear_approximation. collocation.time_dependent_collocation_scheme.TimeDiscreteCollocationScheme` uses here (``dU/dt + A(U) = f``, one derivative, no ``d_t()`` anywhere in the residual). Same reason this directory's ``2d_wave_equation_system.py`` doesn't reuse the library's ``WaveResidual`` either -- see that file's docstring. :func:`NonlinearReactionDiffusionResidual.construct_residual` below is the direct nonlinear analogue of the library's own ``LaplacianResidual`` (``-lap``, reused as-is by ``1d_heat_equation.py``): just ``-lap + u^2``, written as ``rho * rho`` and never ``rho ** 2`` (CLAUDE.md: an exponent on a vmapped leaf differentiates through :math:`0^0` to ``NaN`` -- the same pattern as the ``u1 * u2 * u2`` in ``2d_gray_scott_reaction_diffusion .py``'s weak form, here with both factors the same variable). **Time integrator: Crank-Nicolson**, not Gray-Scott's L-stable Pareschi-Russo: the reaction here is mild (:math:`u^2` with :math:`u = O(1)`, no stiff fast/slow split), so nothing demands an SDIRK -- the same non-stiff reasoning ``2d_wave_equation_system.py`` gives for its own choice of Crank-Nicolson. The problem being nonlinear (unlike that wave system) does not change this: per CLAUDE.md ("le solve Galerkin est deja un Newton... un schema lineaire est le cas ou une iteration suffit"), every stage is a Newton solve regardless of the tableau, so an implicit, non-dissipative tableau costs nothing extra in kind, only (slightly) in Newton iterations. **Kernel space: same conditioning knowledge as ``1d_heat_equation.py``**, reused rather than re-derived, since the spatial operator collocated here still contains that same Laplacian (``sigma_factor`` around 2-3 needed for a well-behaved dense Newton solve; ``sigma_factor = 3.0`` at ``n_points = 20`` below, checked to remain well-conditioned with the reaction term added). """ import time import jax import jax.numpy as jnp import matplotlib.pyplot as plt import numpy as np from scimba_jax.domains.meshless_domains.domains_1d import Segment1D from scimba_jax.linear_approximation.basis.kernel_basis import KernelBasis from scimba_jax.linear_approximation.basis.kernel_function import GaussianKernel from scimba_jax.linear_approximation.collocation.time_discrete_collocation_scheme import ( TimeDiscreteCollocationScheme, ) from scimba_jax.linear_approximation.variables.collocation_variables import ( CollocationVariables, ) from scimba_jax.nonlinear_approximation.model_class.funcparam_vectorial import ( ParamScalarFunction, ) from scimba_jax.physical_models.abstract_residuals import ( NDARRAYS_FUNC_TYPE, PARAM_FUNC_TYPE, InteriorResidual, ) from scimba_jax.time_discrete.butcher_tableau import build_crank_nicolson_tableau # ── Manufactured solution and its forcing term ──────────────────────────────── def u0(x, mu=None): """Initial condition u(x, 0) = sin(pi x).""" return jnp.sin(jnp.pi * x) def u_exact(t, x): """Exact solution u(x, t) = exp(-t) sin(pi x).""" return jnp.exp(-t) * jnp.sin(jnp.pi * x) def f_source(x: jnp.ndarray, mu, t: float) -> jnp.ndarray: """RHS of du/dt + A(u) = f, with A(u) = -u_xx + u^2 (see module docstring). With u = exp(-t) sin(pi x): du/dt = -u -u_xx = pi^2 u u^2 = exp(-2t) sin(pi x)^2 """ u = jnp.exp(-t) * jnp.sin(jnp.pi * x) du_dt = -u diffusion_term = (jnp.pi**2) * u reaction_term = u * u return du_dt + diffusion_term + reaction_term # ── Residual: dU/dt + A(U) = f, A(u) = -u_xx + u^2 ──────────────────────────── class NonlinearReactionDiffusionResidual(InteriorResidual): r"""``A(u) = -u_xx + u^2`` such that :math:`\partial_t u + A(u) = f`. Written directly (no bilinear/linear split: collocation is strong-form, test-function-free), the nonlinear analogue of the library's ``LaplacianResidual`` (``-lap``, reused unchanged by ``1d_heat_equation.py``): the reaction term is the only addition, and it is written ``rho * rho``, never ``rho ** 2`` (see module docstring). """ def __init__( self, domain, f_rhs: NDARRAYS_FUNC_TYPE | None = None, model_type: str = "x_mu" ): super().__init__(domain=domain, size=1, model_type=model_type, f_rhs=f_rhs) def construct_residual(self, *vars: PARAM_FUNC_TYPE) -> PARAM_FUNC_TYPE: rho = vars[0] assert isinstance(rho, ParamScalarFunction) lap = rho.laplacian("x") return -lap + rho * rho # ── Kernel collocation space and solve ──────────────────────────────────────── def make_variables(n_points: int, sigma_factor: float = 3.0): """Build a Gaussian-kernel collocation space on ``n_points`` centers over [0, 1]. Same construction as ``1d_heat_equation.py``: the kernel width is tied to the point spacing (``sigma = sigma_factor * spacing``). """ spacing = 1.0 / (n_points - 1) sigma = sigma_factor * spacing x_centers = jnp.linspace(0.0, 1.0, n_points).reshape(-1, 1) x_centers_bc = jnp.array([[0.0], [1.0]]) domain = Segment1D(jnp.array([[0.0, 1.0]]), is_main_domain=True) basis = KernelBasis( dim=1, output_dim=1, kernel_function=GaussianKernel(sigma=sigma), centers=x_centers, basis_type="scalar", ) variables = CollocationVariables(basis=basis, nb_variables=1) return variables, domain, x_centers, x_centers_bc def solve(n_points: int, dt: float, nt: int, sigma_factor: float = 3.0): variables, domain, x_centers, x_centers_bc = make_variables(n_points, sigma_factor) spatial_residual_factory = lambda t: NonlinearReactionDiffusionResidual( # noqa: E731 domain=domain, f_rhs=lambda x, mu: f_source(x, mu, t), ) scheme = TimeDiscreteCollocationScheme( spatial_residual_factory=spatial_residual_factory, variables=variables, collocation_points=x_centers, bc_collocation_points=x_centers_bc, main_domain=domain, butcher_tableau=build_crank_nicolson_tableau(), dt=dt, dirichlet_factory=lambda t: (lambda x, n, mu: jnp.zeros(1)), max_iter=20, tol=1e-10, ) dofsl_init = scheme.initialize(u0) dofsl_final, history = scheme.solve(dofsl_init, t0=0.0, nt=nt) return variables, dofsl_final, history # ── Evaluation against the manufactured exact solution ──────────────────────── def profile(variables, dofsl, xs): """u_h(xs) for a given DOF array, xs shape (n,) -> values shape (n,).""" variables.dofsl = dofsl return np.asarray( jax.vmap(lambda x: variables.evaluate(x[None]))(jnp.asarray(xs)) ).reshape(-1) def u_exact_profile(t, xs): """u_exact(t, .) sampled at xs, xs shape (n,) -> values shape (n,).""" return np.asarray(jax.vmap(lambda x: u_exact(t, x))(jnp.asarray(xs))) if __name__ == "__main__": N_POINTS = 20 SIGMA_FACTOR = 2.0 NT = 100 T_FINAL = 0.1 DT = T_FINAL / NT print( f"Solving the nonlinear reaction-diffusion equation with kernel collocation: " f"n_points={N_POINTS}, sigma_factor={SIGMA_FACTOR}, " f"dt={DT:.1e}, nt={NT} (T_final={T_FINAL:.4f}) -- Crank-Nicolson, Dirichlet" ) t0 = time.time() variables, dofsl_final, _history = solve(N_POINTS, DT, NT, SIGMA_FACTOR) # block_until_ready: the solve is dispatched asynchronously, so a plain # time.time() delta here would just measure dispatch, not compute. jax.block_until_ready(dofsl_final) print(f"Collocation solve done in {time.time() - t0:.1f}s") n_eval = 400 xs_eval = jnp.linspace(0.0, 1.0, n_eval) u_h = profile(variables, dofsl_final, xs_eval) u_ex = u_exact_profile(T_FINAL, xs_eval) err = np.abs(u_h - u_ex) rel_l2 = np.linalg.norm(err) / np.linalg.norm(u_ex) print(f"rel. L2 error vs. exact = {rel_l2:.3e}, max abs = {err.max():.3e}") # ── Plot: collocation vs. exact profile, and the pointwise error ───────── fig, (ax_u, ax_err) = plt.subplots(1, 2, figsize=(11, 4.5)) xs_np = np.array(xs_eval) ax_u.plot(xs_np, u_ex, "-", label="exact", linewidth=2) ax_u.plot(xs_np, u_h, "--", label="collocation") ax_u.set_xlabel("x") ax_u.set_ylabel("u") ax_u.set_title(f"$u(x, t={T_FINAL:.3f})$") ax_u.legend() ax_err.plot(xs_np, err) ax_err.set_xlabel("x") ax_err.set_ylabel(r"$|u_h - u_\mathrm{exact}|$") ax_err.set_title(f"pointwise error (rel. L2 = {rel_l2:.2e})") fig.suptitle( r"Nonlinear reaction-diffusion, $\partial_t u - \partial_{xx} u + u^2 = f$, " f"kernel collocation (Crank-Nicolson, {N_POINTS} points, Dirichlet)" ) fig.tight_layout() plt.show()