r"""Shallow water over a bump, DG basis enriched with a PINN prior. The well-balanced strategy of ``1d_linear_advection.py`` (Franck, Michel-Dansac, Navoret, "Approximately well-balanced Discontinuous Galerkin methods using bases enriched with Physics-Informed Neural Networks", arXiv:2310.14754), applied to the shallow water *system*: .. math:: \partial_t h + \partial_x q = 0, \qquad \partial_t q + \partial_x \left(\frac{q^2}{h} + \frac{1}{2} g h^2\right) = -g h \, Z'(x), over a Gaussian bump :math:`Z(x) = z_{max} \exp(-(x - 1/2)^2 / (2 \sigma^2))`. A *moving* steady state has :math:`q \equiv q_0` and .. math:: \frac{q_0^2}{2 h^2} + g (h + Z(x)) = C, a cubic in :math:`h(x)` with two positive roots; the parameters keep the flow subcritical (:math:`\mathrm{Fr} \leq 0.51` at the apex over the whole :math:`(z_{max}, \sigma)` box), so a Newton iteration started at the upstream depth picks the subcritical root (:func:`exact_steady_depth`). A PINN learns :math:`h(x; z_{max}, \sigma)` from the steady residual (``SteadyShallowWater1D``, hard-anchored to ``H_UP`` at both ends) over :math:`(z_{max}, \sigma) \in [0.15, 0.25] \times [0.08, 0.12]`; :math:`q`'s prior is the exact constant :math:`q_0`. Frozen at the canonical bump :math:`(0.2, 0.1)`, the prior :math:`(h, q_0)` enriches the Taylor basis MULTIPLICATIVELY (every mode times the prior; ``"with_prior_multiplicative"`` of ``enriched_basis.py``). Both components stay positive, so nothing degenerates; the additive enrichment (the last mode replaced by the prior) would, at degree 1, replace :math:`q`'s linear mode by a second constant. Started at the discrete steady state, a plain Taylor basis drifts away from it -- the discrete analogue of the C-property failure of non-well-balanced schemes -- and :math:`q` drifts with :math:`h`: its flux and source both depend on :math:`h`, so an :math:`h` that the basis cannot represent leaves an imbalance that forces :math:`q` too. The enriched basis nearly contains the steady state, and the drift is what is left of the prior's error. The comparisons of the former version of this file (plain / additive / multiplicative bases, and a 50-bump parametric study solved in one ``jax.vmap``), with those of the former ``1d_shallow_water_bases_and_integrators.py`` (Lagrange basis, DIRK(4,5), a perturbed steady state), live in the benchmark ``benchmarks/benchmarks_jax/dg_enriched_well_balanced/`` (``shallow_water_1d`` dataset). """ # %% import time 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.shallow_water_weak_form import ( ShallowWaterWeakForm1D, make_shallow_water_flux_fns, ) from scimba_jax.physical_models.temporal_pde.shallow_water import SteadyShallowWater1D from scimba_jax.plots.plots_nd import plot_abstract_approx_spaces from scimba_jax.time_discrete.butcher_tableau import build_rk4_tableau DIM = 1 OUT_DIM = 2 # (h, q) X_MIN, X_MAX = 0.0, 1.0 GRAVITY = 1.0 H_UP = 1.0 # upstream depth, at x = X_MIN (where the bump is negligible) Q0 = 0.3 # constant discharge of the moving steady state; Fr(H_UP) = 0.3 X_C = 0.5 # bump center, fixed (not a physical parameter) # The PINN prior is trained over the whole (z_max, sigma) box; the DG solve # runs at the canonical bump (Z_MAX_DEFAULT, SIGMA_DEFAULT). Z_MAX_MIN, Z_MAX_MAX = 0.15, 0.25 SIGMA_MIN, SIGMA_MAX = 0.08, 0.12 DOM_MU = [[Z_MAX_MIN, Z_MAX_MAX], [SIGMA_MIN, SIGMA_MAX]] Z_MAX_DEFAULT, SIGMA_DEFAULT = 0.2, 0.1 MU_DEFAULT = jnp.array([Z_MAX_DEFAULT, SIGMA_DEFAULT]) CATEGORY = "with_prior_multiplicative" N_CELLS = 40 ORDER = 1 T_FINAL = 1.0 DEFAULT_HIDDEN_SIZES = [16] * 2 # ── Bathymetry and the exact steady state ──────────────────────────────────── def bathymetry(x: jnp.ndarray, z_max: jnp.ndarray, sigma: jnp.ndarray) -> jnp.ndarray: """Gaussian bump ``Z(x; z_max, sigma)``, elementwise on ``x``. Handed as ``bathymetry=`` to the PINN residual and to the DG weak form, which both differentiate it internally: no hand-derived ``Z'``. """ dx = x - X_C return z_max * jnp.exp(-(dx * dx) / (2.0 * sigma * sigma)) def steady_state_const(z_max: jnp.ndarray, sigma: jnp.ndarray) -> jnp.ndarray: """The constant ``C``, fixed by ``h = H_UP`` at ``x = X_MIN`` (``Z`` included).""" z_upstream = bathymetry(jnp.array([X_MIN]), z_max, sigma)[0] return 0.5 * Q0 * Q0 / (H_UP * H_UP) + GRAVITY * (H_UP + z_upstream) def exact_steady_depth( x: jnp.ndarray, z_max: jnp.ndarray, sigma: jnp.ndarray, n_iter: int = 30 ) -> jnp.ndarray: r"""The subcritical root of ``g h^3 + (g Z(x) - C) h^2 + q0^2/2 = 0``. Newton from ``H_UP``: the two positive roots stay well separated over the parameter box (``Fr <= 0.51`` at the apex), so the iteration never crosses into the supercritical root's basin. Unrolled (``n_iter`` is a static int), so it traces, vmaps and differentiates like a closed form. Args: x: physical coordinate(s), shape ``(1,)`` (or any batchable shape). z_max: bump height (physical parameter). sigma: bump width (physical parameter). n_iter: Newton iterations (far more than double precision needs). Returns: ``h(x)``, same shape as ``x``. """ z = bathymetry(x, z_max, sigma) c = GRAVITY * z - steady_state_const(z_max, sigma) h = jnp.full_like(x, H_UP) for _ in range(n_iter): h2 = h * h residual = GRAVITY * h2 * h + c * h2 + 0.5 * Q0 * Q0 residual_prime = 3.0 * GRAVITY * h2 + 2.0 * c * h h = h - residual / residual_prime return h def exact_steady_depth_wrt_mu(x: jnp.ndarray, mu: jnp.ndarray) -> jnp.ndarray: """``h(x; mu)`` on a batch of points and parameters (for the prior's plot).""" z_max, sigma = mu[:, 0, None], mu[:, 1, None] steady = jax.vmap(exact_steady_depth, in_axes=(0, 0, 0)) return steady(x, z_max, sigma) def steady_state( x: jnp.ndarray, z_max: jnp.ndarray = Z_MAX_DEFAULT, sigma: jnp.ndarray = SIGMA_DEFAULT, ) -> jnp.ndarray: """The exact steady state ``(h(x), q0)``, shape ``(2,)``, at the canonical bump. Initial state and Dirichlet datum of the DG solve. """ return jnp.concatenate([exact_steady_depth(x, z_max, sigma), jnp.array([Q0])]) # ── PINN: steady-state prior for h(x) alone (q's prior is the constant q0) ── def post_processing( inputs: jnp.ndarray, x: jnp.ndarray, mu: jnp.ndarray ) -> jnp.ndarray: r"""Hard boundary condition: ``h = H_UP`` EXACTLY at both ends. ``X_C`` is the midpoint of ``[X_MIN, X_MAX]``, so the steady-state cubic is the same at both ends and its root is ``H_UP`` at both. The network output is weighted by ``t (1 - t)``, ``t = (x - X_MIN) / (X_MAX - X_MIN)``, which vanishes at both ends, and exponentiated, so ``h > 0``. ``mu`` is unused but part of the ``model_type="x_mu"`` signature. """ t = (x[0] - X_MIN) / (X_MAX - X_MIN) bump = t * (1.0 - t) return H_UP * jnp.exp(bump * inputs) def _bathymetry_x_mu(x: jnp.ndarray, mu: jnp.ndarray) -> jnp.ndarray: """``Z(x; z_max, sigma)``, ``mu = (z_max, sigma)``, for the PINN residual.""" return bathymetry(x, mu[0:1], mu[1:2]) def train_prior( key: jax.Array, hidden_sizes: tuple[int, ...] = DEFAULT_HIDDEN_SIZES, n_epochs: int = 1000, n_colloc: int = 2000, optimizer: str = "ENG", ) -> tuple[jax.Array, Projector]: """Trains the steady-state PINN prior for ``h(x; mu)`` over the whole 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 = [(Z_MAX_MIN, Z_MAX_MAX), (SIGMA_MIN, SIGMA_MAX)] sampler = TensorizedSampler( [DomainSampler(domain_x), UniformParametricSampler(domain_mu)], bc=False ) key, key_mlp = jax.random.split(key) nn = MLP( in_size=3, out_size=1, hidden_sizes=list(hidden_sizes), activation="tanh", key=key_mlp, ) space = ApproximationSpace( {"x": 1, "mu": 2}, [(nn, "scalar", 1)], model_type="x_mu", post_processing=post_processing, ) model = SteadyShallowWater1D( domain_x, gravity=GRAVITY, discharge=Q0, bathymetry=_bathymetry_x_mu, model_type="x_mu", ) pinn = Projector( model, space, sampler, optimizer=optimizer, adaptive_matrix_regularization=True, linesearch="armijo", ) return pinn.project(key, space, n_epochs, n_colloc) def append_discharge(x: jnp.ndarray, h: jnp.ndarray) -> jnp.ndarray: """``(h, q0)``: the exact constant discharge appended to the ``h`` prior.""" return jnp.concatenate([h, jnp.array([Q0])]) # ── DG: shallow water weak form and flux ───────────────────────────────────── def make_scheme(variables, dt: float) -> TimeDiscreteDGscheme: """RK4 in time, Rusanov flux, Dirichlet datum at the steady state.""" return TimeDiscreteDGscheme( spatial_weak_form_factory=ShallowWaterWeakForm1D( dim=DIM, gravity=GRAVITY, bathymetry=bathymetry, args=(MU_DEFAULT[0:1], MU_DEFAULT[1:2]), ), variables=variables, flux=LocalLaxFriedrichsFlux(*make_shallow_water_flux_fns(GRAVITY)), butcher_tableau=build_rk4_tableau(), dt=dt, dirichlet=steady_state, ) def plot_results(variables, dofsl_init, dofsl_final, errors, x_plot): """Errors of ``h`` and ``q`` at ``t = 0`` and ``T_FINAL``, and the free surface.""" fig, axes = plt.subplots(1, 3, figsize=(15, 5)) w_ex = jax.vmap(steady_state)(x_plot) z = jax.vmap(bathymetry, in_axes=(0, None, None))( x_plot, Z_MAX_DEFAULT, SIGMA_DEFAULT )[:, 0] for label, dofsl in (("t=0", dofsl_init), (f"t={T_FINAL}", dofsl_final)): variables.dofsl = dofsl w_h = variables.evaluate(x_plot) axes[0].semilogy(x_plot[:, 0], jnp.abs(w_h[:, 0] - w_ex[:, 0]), label=label) axes[1].semilogy(x_plot[:, 0], jnp.abs(w_h[:, 1] - w_ex[:, 1]), label=label) for ax, name, i in ((axes[0], "h", 0), (axes[1], "q", 1)): ax.set_xlabel("x") ax.set_ylabel(f"|{name}_h - {name}|") ax.set_title( f"{name}: relative L2 distance {errors[0][i]:.2e} (t=0), " f"{errors[1][i]:.2e} (t={T_FINAL})" ) ax.legend(fontsize=8) ax.grid(True, alpha=0.3) variables.dofsl = dofsl_final w_h = variables.evaluate(x_plot) axes[2].plot(x_plot[:, 0], w_ex[:, 0] + z, "k--", label="exact h + Z") axes[2].plot(x_plot[:, 0], w_h[:, 0] + z, label=f"enriched DG, t={T_FINAL}") axes[2].plot(x_plot[:, 0], z, "gray", label="bathymetry Z") axes[2].set_xlabel("x") axes[2].set_ylabel("h + Z") axes[2].set_title("Free surface") axes[2].legend(fontsize=8) axes[2].grid(True, alpha=0.3) fig.suptitle( f"Shallow water over a bump, moving steady state, {CATEGORY}\n" f"q0={Q0}, g={GRAVITY}, H_up={H_UP}, z_max={Z_MAX_DEFAULT}, " f"sigma={SIGMA_DEFAULT}, n_cells={N_CELLS}, order={ORDER}" ) fig.tight_layout() plt.show() # %% if __name__ == "__main__": key = jax.random.PRNGKey(0) print("training the PINN prior for h(x; z_max, sigma) over the parameter box ...") _t0 = time.perf_counter() key, pinn = train_prior(key) jax.block_until_ready((key, pinn.best_loss)) print( f"PINN prior trained in {time.perf_counter() - _t0:.2f} s, " f"best loss = {pinn.best_loss}" ) plot_abstract_approx_spaces( [pinn.space], pinn.model.main_domain, DOM_MU, loss=pinn.losses, components=[{"height": 0}], solution=exact_steady_depth_wrt_mu, error=exact_steady_depth_wrt_mu, title="PINN prior on the steady depth h, vs. the Newton solution.", ) plt.show() prior_fn = make_prior_fn(pinn, MU_DEFAULT, postprocess=append_discharge) h = 1.0 / N_CELLS speed_scale = 1.5 # max |u| + sqrt(g h) over the parameter box (~1.3) dt = 0.2 * h / ((2 * ORDER + 1) * speed_scale) nt = round(T_FINAL / dt) mesh = cartesian_mesh(n_cells=[N_CELLS], quad_order=ORDER + 3) variables = make_variables(mesh, ORDER, OUT_DIM, 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: " f"h {err_init[0]:.3e} / q {err_init[1]:.3e} (t=0), " f"h {err_final[0]:.3e} / q {err_final[1]:.3e} (t={T_FINAL})" ) plot_results(variables, dofsl_init, dofsl_final, (err_init, err_final), x_plot) # %%