r"""Linear transport with a reaction source, DG basis enriched with a PINN prior. Reproduces the well-balanced test case of Franck, Michel-Dansac, Navoret, "Approximately well-balanced Discontinuous Galerkin methods using bases enriched with Physics-Informed Neural Networks", arXiv:2310.14754 (https://github.com/Victor-MichelDansac/DG-PINNs): .. math:: u_t + c u_x = a u + b u^2 \quad \text{on } x \in (0, 1), \qquad u(t, 0) = u_0, with :math:`c = 1`, whose steady solution is :math:`u(x) = a u_0 / ((a + b u_0) e^{-a x} - b u_0)`. A PINN is trained once on the STEADY equation over a box of parameters :math:`\mu = (a, b, u_0)` (``SteadyAdvectionReactionND``, the boundary value hard-constrained), then frozen at :math:`\mu = (0.75, 0.75, 0.15)`. Its output :math:`g` enriches the per-cell Taylor basis MULTIPLICATIVELY, :math:`\varphi_k = T_k \, g` (``"with_prior_multiplicative"`` of ``enriched_basis.py``): every mode carries the prior, so the local space contains :math:`g` times every polynomial of degree ``ORDER``, and the steady state, which is close to :math:`g`, is nearly representable. The choice is safe here because :math:`u` (hence its prior) stays positive over the whole parameter box; the additive enrichment (the last Taylor mode replaced by :math:`g`) and the plain basis are compared in the benchmark. The scheme is started at the discrete steady state and advanced to ``T_FINAL``: how far it drifts away measures how well balanced it is. A plain polynomial basis cannot represent the steady state, so its flux and its source do not cancel and the solution moves; the enriched basis nearly contains it, and the drift is what is left of the prior's error. Fed the exact steady state instead of the PINN, the enriched scheme is exactly well balanced. The comparisons of the former version of this file (plain / additive / multiplicative bases, the exact-prior oracle, and the spatial order of every basis against an unsteady exact solution) live in the benchmark ``benchmarks/benchmarks_jax/dg_enriched_well_balanced/`` (``linear_advection`` dataset). """ # %% import jax import jax.numpy as jnp import matplotlib.pyplot as plt from scimba_jax.domains.meshless_domains.domains_1d import Segment1D from scimba_jax.linear_approximation.galerkin.dg.enriched_approach.enriched_basis import ( # noqa: E501 make_variables, ) from scimba_jax.linear_approximation.galerkin.dg.enriched_approach.well_balanced_enrichment import ( # noqa: E501 make_prior_fn, relative_l2_errors, solve, ) from scimba_jax.linear_approximation.galerkin.dg.flux import LocalLaxFriedrichsFlux from scimba_jax.linear_approximation.galerkin.dg.time_discrete_dg_scheme import ( TimeDiscreteDGscheme, ) from scimba_jax.linear_approximation.meshes.cartesian_mesh import cartesian_mesh from scimba_jax.nonlinear_approximation.approximation_spaces.approximation_spaces import ( # noqa: E501 ApproximationSpace, ) from scimba_jax.nonlinear_approximation.integration.monte_carlo import ( DomainSampler, TensorizedSampler, ) from scimba_jax.nonlinear_approximation.integration.monte_carlo_parameters import ( UniformParametricSampler, ) from scimba_jax.nonlinear_approximation.networks.mlp import MLP from scimba_jax.nonlinear_approximation.numerical_solvers.projectors import Projector from scimba_jax.physical_models.classical_weakform.advection_reaction_weak_form import ( # noqa: E501 AdvectionReactionWeakForm, ) from scimba_jax.physical_models.temporal_pde.advection import SteadyAdvectionReactionND from scimba_jax.time_discrete.butcher_tableau import build_rk4_tableau DIM = 1 X_MIN, X_MAX = 0.0, 1.0 A_MIN, A_MAX = 0.5, 1.0 B_MIN, B_MAX = 0.5, 1.0 U0_MIN, U0_MAX = 0.1, 0.2 SPEED = 1.0 REACTION_A, REACTION_B, LEFT_BC = 0.75, 0.75, 0.15 MU_DEFAULT = jnp.array([REACTION_A, REACTION_B, LEFT_BC]) CATEGORY = "with_prior_multiplicative" N_CELLS = 40 ORDER = 1 T_FINAL = 0.5 DEFAULT_HIDDEN_SIZES = [16] * 2 def steady_solution(x: jnp.ndarray, a: float, b: float, u0: float) -> jnp.ndarray: """Closed-form steady solution of ``u' = a u + b u^2``, ``u(0) = u0``. Elementwise in ``x``: a single point (shape ``(1,)``) or a batch. """ return a * u0 / ((a + b * u0) * jnp.exp(-a * x) - b * u0) def steady_state(x: jnp.ndarray) -> jnp.ndarray: """The steady state at ``MU_DEFAULT``: initial state and Dirichlet datum.""" return steady_solution(x, *MU_DEFAULT) # ── PINN: steady-state prior ───────────────────────────────────────────────── def post_processing( inputs: jnp.ndarray, x: jnp.ndarray, mu: jnp.ndarray ) -> jnp.ndarray: """``u0 + x * NN(x, a, b, u0)``: hard-constrains ``u(0) = u0`` exactly.""" u0 = mu[2:3] return u0 + x * inputs def _coefficient_a(x, mu): return mu[0:1] def _coefficient_b(x, mu): return mu[1:2] def reaction(u): """``a u + b u^2``, coefficients read straight off ``mu`` (model_type "x_mu").""" return u * _coefficient_a + (u * u) * _coefficient_b def train_prior( key: jax.Array, hidden_sizes: tuple[int, ...] = DEFAULT_HIDDEN_SIZES, n_epochs: int = 500, n_colloc: int = 2000, optimizer: str = "ENG", ) -> tuple[jax.Array, Projector]: """Trains the steady-state PINN prior over the whole parameter box. Args: key: PRNG state. hidden_sizes: MLP hidden layer sizes. n_epochs: number of optimizer epochs. n_colloc: number of collocation points per epoch. optimizer: scimba optimizer name (natural gradient "ENG" by default). Returns: The updated PRNG key and the trained :class:`Projector`. """ domain_x = Segment1D((X_MIN, X_MAX), is_main_domain=True) domain_mu = [(A_MIN, A_MAX), (B_MIN, B_MAX), (U0_MIN, U0_MAX)] sampler = TensorizedSampler( [DomainSampler(domain_x), UniformParametricSampler(domain_mu)], bc=False ) key, key_mlp = jax.random.split(key) nn = MLP( in_size=4, out_size=1, hidden_sizes=list(hidden_sizes), activation="tanh", key=key_mlp, ) space = ApproximationSpace( {"x": 1, "mu": 3}, [(nn, "scalar", 1)], model_type="x_mu", post_processing=post_processing, ) model = SteadyAdvectionReactionND( domain_x, time_domain=None, velocity=SPEED, reaction=reaction, model_type="x_mu", bc=None, ) pinn = Projector( model, space, sampler, optimizer=optimizer, adaptive_matrix_regularization=True, linesearch="armijo", ) return pinn.project(key, space, n_epochs, n_colloc) # ── DG: balance-law weak form and flux ──────────────────────────────────────── def dg_reaction(u): """``a u + b u^2`` at ``MU_DEFAULT``, the DG side's reaction.""" return REACTION_A * u + REACTION_B * (u * u) def flux_fn(u: jnp.ndarray) -> jnp.ndarray: """``f(u) = c u``.""" return SPEED * u def dflux_fn(u: jnp.ndarray) -> jnp.ndarray: """``f'(u) = c``.""" return SPEED * jnp.ones_like(u) def make_scheme(variables, dt: float) -> TimeDiscreteDGscheme: """RK4 in time, local Lax-Friedrichs flux, Dirichlet datum at the steady state.""" return TimeDiscreteDGscheme( spatial_weak_form_factory=AdvectionReactionWeakForm( dim=DIM, speed=SPEED, reaction=dg_reaction ), variables=variables, flux=LocalLaxFriedrichsFlux(flux_fn, dflux_fn), butcher_tableau=build_rk4_tableau(), dt=dt, dirichlet=steady_state, ) def plot_results(variables, dofsl_init, dofsl_final, errors, x_plot): """The solution at ``T_FINAL`` against the steady state, and the drift.""" fig, axes = plt.subplots(1, 2, figsize=(12, 5)) u_ex = jax.vmap(steady_state)(x_plot)[:, 0] variables.dofsl = dofsl_final u_final = variables.evaluate(x_plot)[:, 0] axes[0].plot(x_plot[:, 0], u_ex, "k--", label="exact steady state") axes[0].plot(x_plot[:, 0], u_final, label=f"enriched DG, t={T_FINAL}") axes[0].set_xlabel("x") axes[0].set_ylabel("u") axes[0].set_title("Solution") axes[0].legend(fontsize=8) axes[0].grid(True, alpha=0.3) for label, dofsl in (("t=0", dofsl_init), (f"t={T_FINAL}", dofsl_final)): variables.dofsl = dofsl u_h = variables.evaluate(x_plot)[:, 0] axes[1].semilogy(x_plot[:, 0], jnp.abs(u_h - u_ex), label=label) axes[1].set_xlabel("x") axes[1].set_ylabel("|u_h - u|") axes[1].set_title( f"Distance to the steady state: {errors[0]:.2e} (t=0), " f"{errors[1]:.2e} (t={T_FINAL})" ) axes[1].legend(fontsize=8) axes[1].grid(True, alpha=0.3) fig.suptitle( "Linear transport with reaction source, u_t + u_x = a u + b u^2, " f"{CATEGORY}\n" f"a={REACTION_A}, b={REACTION_B}, u0={LEFT_BC}, " f"n_cells={N_CELLS}, order={ORDER}" ) fig.tight_layout() plt.show() # %% if __name__ == "__main__": key = jax.random.PRNGKey(0) print("training the PINN prior ...") key, pinn = train_prior(key) print(f"PINN prior trained, best loss = {pinn.best_loss}") prior_fn = make_prior_fn(pinn, MU_DEFAULT) h = 1.0 / N_CELLS dt = 0.2 * h / (2 * ORDER + 1) nt = round(T_FINAL / dt) mesh = cartesian_mesh(n_cells=[N_CELLS], quad_order=ORDER + 3) variables = make_variables(mesh, ORDER, 1, prior_fn, CATEGORY) scheme = make_scheme(variables, dt) dofsl_init, dofsl_final = solve( scheme, steady_state, nt, description=f"DG solve ({CATEGORY})" ) x_plot = jnp.linspace(X_MIN, X_MAX, 400)[:, None] (err_init,) = relative_l2_errors(variables, dofsl_init, steady_state, x_plot) (err_final,) = relative_l2_errors(variables, dofsl_final, steady_state, x_plot) print( f"relative L2 distance to the steady state: {err_init:.3e} (t=0), " f"{err_final:.3e} (t={T_FINAL})" ) plot_results(variables, dofsl_init, dofsl_final, (err_init, err_final), x_plot) # %%