r"""2D shallow water over an axisymmetric bump, DG basis enriched with a PINN prior. The 2D case of Franck, Michel-Dansac, Navoret, "Approximately well-balanced Discontinuous Galerkin methods using bases enriched with Physics-Informed Neural Networks", arXiv:2310.14754, on the square :math:`[-3, 3]^2`: .. math:: \partial_t h + \nabla \cdot Q = 0, \qquad \partial_t Q + \nabla \cdot \left(\frac{Q \otimes Q}{h} + \frac{1}{2} g h^2 \, \mathrm{Id}\right) = -g h \, \nabla Z(x; \mu), with :math:`Z(x; \mu) = \gamma \exp(\alpha (r_0^2 - r^2))`, :math:`\mu = (\alpha, \gamma, r_0)`. The steady state is a rotating vortex in closed form, .. math:: h = 2 - Z - \frac{\alpha}{8 g} Z^4, \qquad Q = (x_2, -x_1) \, h \, u, \quad u = \alpha Z^2, divergence-free for ANY radial :math:`h u` (so the mass equation holds by construction). **The prior.** The PINN predicts the two scalars :math:`(h, u)` over the box :math:`\mu \in [0.25, 0.75] \times [0.1, 0.4] \times [0.5, 1.25]`, trained on the steady radial momentum balance :math:`g \nabla(h + Z) = u^2 x` (``SteadyShallowWater2D`` with the swirl ``components``); :math:`Q` is rebuilt as :math:`(x_2, -x_1) h u`. That residual alone leaves :math:`h`'s constant and :math:`u`'s shape free (trained without an anchor, the network converged to :math:`u \approx 0` and an arbitrary constant :math:`h`), so :func:`post_processing` anchors :math:`(h, u)` exactly to the closed form at :math:`r = 0` and :math:`r = L` (:math:`u` in log space, so it stays positive). Predicting :math:`(h, u)` rather than a stream function keeps the basis gradient a FIRST derivative of the network: a stream-function prior roughly doubled the prior's compile-time overhead. **The basis.** ``"with_prior_mixed"``: multiplicative on :math:`h` (positive, near ``H0``), additive on :math:`q_x, q_y` (the last Taylor mode replaced by the prior). The discharge components cross zero on the vortex's symmetry axes, where a multiplicative prior would collapse their local basis (rank-deficient mass matrix). The Krylov solve of each RK stage is capped at one warm-started iteration (``MAX_ITER_LINEAR``; no measurable difference against 64). **The lesson.** Started at the discrete steady state, a plain Taylor basis drifts away from it (flux and source do not cancel on a non-polynomial state); the enriched basis nearly contains the vortex, and the drift at ``T_FINAL`` is what is left of the prior's error. The comparisons of the former version of this file (plain vs enriched basis, a 30-vortex parametric study in one ``jax.vmap``), and the disk version of this case (curved unstructured mesh, perturbed steady state) live in the benchmark ``benchmarks/benchmarks_jax/dg_enriched_well_balanced/`` (``shallow_water_2d`` dataset). """ import time import jax import jax.numpy as jnp import matplotlib.pyplot as plt from scimba_jax.domains.meshless_domains.domains_2d import Square2D 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 ( ShallowWater2DWeakForm, make_shallow_water_flux_fns_2d, ) from scimba_jax.physical_models.temporal_pde.shallow_water import SteadyShallowWater2D from scimba_jax.plots.plots_nd import plot_abstract_approx_spaces from scimba_jax.time_discrete.butcher_tableau import build_rk4_tableau DIM = 2 OUT_DIM = 3 # (h, qx, qy) GRAVITY = 1.0 H0 = 2.0 # background depth, far from the bump (Z -> 0 there) L = 3.0 # domain is [-L, L]^2 # The PINN prior is trained over this whole box; the DG solve runs at the # midpoint (ALPHA_DEFAULT, GAMMA_DEFAULT, R0_DEFAULT). ALPHA_MIN, ALPHA_MAX = 0.25, 0.75 GAMMA_MIN, GAMMA_MAX = 0.1, 0.4 R0_MIN, R0_MAX = 0.5, 1.25 DOM_MU = [[ALPHA_MIN, ALPHA_MAX], [GAMMA_MIN, GAMMA_MAX], [R0_MIN, R0_MAX]] ALPHA_DEFAULT, GAMMA_DEFAULT, R0_DEFAULT = 0.5, 0.25, 0.875 MU_DEFAULT = jnp.array([ALPHA_DEFAULT, GAMMA_DEFAULT, R0_DEFAULT]) CATEGORY = "with_prior_mixed" MULTIPLICATIVE_INDICES = (0,) # h; qx and qy are enriched additively N_CELLS = (24, 24) ORDER = 1 MAX_ITER_LINEAR = 1 T_FINAL = 1.0 N_EPOCHS = 1000 DEFAULT_HIDDEN_SIZES = [16] * 3 BOX_BOUNDS = ((-L, L), (-L, L)) # ── Bathymetry and the steady state ────────────────────────────────────────── def bathymetry(x: jnp.ndarray, alpha: jnp.ndarray, gamma: jnp.ndarray, r0: jnp.ndarray): """``Z(x; alpha, gamma, r0) = gamma * exp(alpha * (r0^2 - r^2))``, ``x`` of shape ``(2,)``. Handed as ``bathymetry=`` to the PINN residual and to the DG weak form, which both differentiate it internally: no hand-derived ``grad Z``. """ r2 = x[0] * x[0] + x[1] * x[1] return gamma * jnp.exp(alpha * (r0 * r0 - r2)) def steady_h_and_u( x: jnp.ndarray, alpha: jnp.ndarray, gamma: jnp.ndarray, r0: jnp.ndarray ) -> tuple[jnp.ndarray, jnp.ndarray]: """The depth ``h`` and the swirl speed ``u`` such that ``Q = (x2, -x1) h u``.""" z = bathymetry(x, alpha, gamma, r0) z2 = z * z z4 = z2 * z2 h = H0 - z - (alpha / (8.0 * GRAVITY)) * z4 u = alpha * z2 return h, u def steady_h_and_u_wrt_mu(x: jnp.ndarray, mu: jnp.ndarray) -> jnp.ndarray: """``(h, u)`` on a batch of points and parameters (for the prior's plot).""" alpha, gamma, r0 = mu[:, 0, None], mu[:, 1, None], mu[:, 2, None] steady = jax.vmap(steady_h_and_u, in_axes=(0, 0, 0, 0)) return jnp.concatenate(steady(x, alpha, gamma, r0), axis=-1) def steady_state( x: jnp.ndarray, alpha: jnp.ndarray = ALPHA_DEFAULT, gamma: jnp.ndarray = GAMMA_DEFAULT, r0: jnp.ndarray = R0_DEFAULT, ) -> jnp.ndarray: """The steady state ``(h, qx, qy)``, shape ``(3,)``, at the default vortex. Initial state and Dirichlet datum of the DG solve. """ h, u = steady_h_and_u(x, alpha, gamma, r0) qx = h * u * x[1] qy = -h * u * x[0] return jnp.array([h, qx, qy]) def default_bathymetry(x: jnp.ndarray) -> jnp.ndarray: """``Z`` at the default vortex (for the plots).""" return bathymetry(x, ALPHA_DEFAULT, GAMMA_DEFAULT, R0_DEFAULT) # ── PINN: (h, u) prior over mu = (alpha, gamma, r0) ───────────────────────── def post_processing( inputs: jnp.ndarray, x: jnp.ndarray, mu: jnp.ndarray ) -> jnp.ndarray: r"""Hard boundary condition: ``(h, u)`` EXACT at ``r = 0`` and ``r = L``. A linear interpolation in ``r^2`` between the two closed-form anchors, plus the network output times ``r2_0 * r2_L`` (zero at both). ``u`` is interpolated and corrected in log space, then exponentiated: it stays positive and the correction stays two-sided (squaring it, an earlier version, only allowed ``u`` above the interpolation and killed the gradient at zero). """ alpha, gamma, r0 = mu[0], mu[1], mu[2] L_ = jnp.full_like(x, L / (2.0**0.5)) h_bc, u_bc = steady_h_and_u(L_, alpha, gamma, r0) center = jnp.full_like(x, 0.0) h_center, u_center = steady_h_and_u(center, alpha, gamma, r0) r2 = x[0] * x[0] + x[1] * x[1] r2_L = 1.0 - r2 / (L * L) r2_0 = r2 / (L * L) h_interp = r2_L * h_center + r2_0 * h_bc log_u_interp = r2_L * jnp.log(u_center) + r2_0 * jnp.log(u_bc) bump = r2_0 * r2_L h = h_interp + bump * inputs[0] u = jnp.exp(log_u_interp + bump * inputs[1]) return jnp.array([h, u]) def _bathymetry_x_mu(x: jnp.ndarray, mu: jnp.ndarray) -> jnp.ndarray: """``Z(x; alpha, gamma, r0)``, ``mu = (alpha, gamma, r0)``, for the PINN residual.""" return bathymetry(x, mu[0], mu[1], mu[2]) def _x_coordinate(x): return x[0] def _y_coordinate(x): return x[1] def _swirl_components(h, u): """``(h, u) -> (h, qx, qy)`` by the swirl ansatz, for the PINN residual.""" qx = h * u * _y_coordinate qy = -h * u * _x_coordinate return h, qx, qy def train_prior( key: jax.Array, hidden_sizes: tuple[int, ...] = DEFAULT_HIDDEN_SIZES, n_epochs: int = N_EPOCHS, n_colloc: int = 4000, optimizer: str = "ENG", file_name: str | None = None, retrain: bool = False, ) -> tuple[jax.Array, Projector]: """Trains the steady-state PINN prior for ``(h, u)(x; alpha, gamma, r0)``. 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). file_name: optional name to save the trained prior under (and load it from on the next run). retrain: train even if a saved prior exists. Returns: The updated PRNG key and the trained :class:`Projector`. """ domain_x = Square2D([(-L, L), (-L, L)], is_main_domain=True) domain_mu = [(ALPHA_MIN, ALPHA_MAX), (GAMMA_MIN, GAMMA_MAX), (R0_MIN, R0_MAX)] sampler = TensorizedSampler( [DomainSampler(domain_x), UniformParametricSampler(domain_mu)], bc=False ) key, key_mlp = jax.random.split(key) nn = MLP( in_size=5, out_size=2, hidden_sizes=list(hidden_sizes), activation="tanh", key=key_mlp, ) space = ApproximationSpace( {"x": 2, "mu": 3}, [(nn, "vec", 2)], model_type="x_mu", post_processing=post_processing, ) model = SteadyShallowWater2D( domain_x, gravity=GRAVITY, bathymetry=_bathymetry_x_mu, components=_swirl_components, model_type="x_mu", ) pinn = Projector( model, space, sampler, optimizer=optimizer, adaptive_matrix_regularization=True, linesearch="armijo", ) if (file_name is not None) and (not retrain): n_epochs_load, pinn = pinn.load(file_name) if n_epochs_load > 0: print(f"Loaded pre-trained model from {file_name}.") return pinn.key, pinn key, pinn = pinn.project(key, space, n_epochs, n_colloc) if file_name is not None: pinn.save(file_name) return key, pinn def _swirl_postprocess(x: jnp.ndarray, h_u: jnp.ndarray) -> jnp.ndarray: """``(h, u) -> (h, qx, qy)`` at a point: the prior the DG basis takes.""" h, u = h_u return jnp.array([h, h * u * x[1], -h * u * x[0]]) # ── DG ─────────────────────────────────────────────────────────────────────── def make_scheme(variables, dt: float) -> TimeDiscreteDGscheme: """RK4, 2D Rusanov flux, Dirichlet datum at the steady state, Krylov capped.""" return TimeDiscreteDGscheme( spatial_weak_form_factory=ShallowWater2DWeakForm( dim=DIM, gravity=GRAVITY, bathymetry=bathymetry, args=(MU_DEFAULT[0], MU_DEFAULT[1], MU_DEFAULT[2]), ), variables=variables, flux=LocalLaxFriedrichsFlux(*make_shallow_water_flux_fns_2d(GRAVITY)), butcher_tableau=build_rk4_tableau(), dt=dt, dirichlet=steady_state, max_iter_linear=MAX_ITER_LINEAR, ) def cfl_dt_nt( n_cells, order: int, t_final: float, speed_scale: float = 1.1 * GRAVITY**0.5 ): """``dt``/``nt`` from an explicit CFL, ``speed_scale`` bounding ``|Q|/h + sqrt(g h)``.""" h_cell = 2.0 * L / min(n_cells) dt = 0.2 * h_cell / ((2 * order + 1) * speed_scale) nt = round(t_final / dt) return dt, nt def error_grid(n_per_axis: int = 50) -> jnp.ndarray: """A structured ``n_per_axis x n_per_axis`` grid over the box.""" grid = jnp.linspace(-L, L, n_per_axis) xx, yy = jnp.meshgrid(grid, grid, indexing="ij") return jnp.stack([xx.ravel(), yy.ravel()], axis=1) # ── Plotting ───────────────────────────────────────────────────────────────── def plot_setup(n_per_axis: int = 120): """The steady vortex: free surface ``h + Z`` and discharge ``Q``.""" grid = jnp.linspace(-L, L, n_per_axis) xx, yy = jnp.meshgrid(grid, grid, indexing="ij") pts = jnp.stack([xx.ravel(), yy.ravel()], axis=1) w_ex = jax.vmap(steady_state)(pts) z_ex = jax.vmap(default_bathymetry)(pts) h_field = (w_ex[:, 0] + z_ex).reshape(n_per_axis, n_per_axis) qx_field = w_ex[:, 1].reshape(n_per_axis, n_per_axis) qy_field = w_ex[:, 2].reshape(n_per_axis, n_per_axis) q_mag = jnp.sqrt(qx_field**2 + qy_field**2) fig, axes = plt.subplots(1, 2, figsize=(11, 5)) pcm0 = axes[0].pcolormesh(xx, yy, h_field, shading="auto", cmap="turbo") fig.colorbar(pcm0, ax=axes[0], shrink=0.85, label="h + Z") axes[0].set_title("steady free surface h + Z") axes[0].set_xlabel("x1") axes[0].set_ylabel("x2") axes[0].set_aspect("equal") pcm1 = axes[1].pcolormesh(xx, yy, q_mag, shading="auto", cmap="viridis") fig.colorbar(pcm1, ax=axes[1], shrink=0.85, label="|Q|") skip = max(1, n_per_axis // 16) axes[1].quiver( xx[::skip, ::skip], yy[::skip, ::skip], qx_field[::skip, ::skip], qy_field[::skip, ::skip], color="white", alpha=0.85, ) axes[1].set_title("steady discharge Q") axes[1].set_xlabel("x1") axes[1].set_ylabel("x2") axes[1].set_aspect("equal") fig.suptitle( f"Steady vortex: alpha={ALPHA_DEFAULT}, gamma={GAMMA_DEFAULT}, " f"r0={R0_DEFAULT}, g={GRAVITY}" ) fig.tight_layout() def plot_results(variables, dofsl_final, errors, n_points: int = 400): """Slice ``x2 = 0`` at ``T_FINAL``: free surface, ``|Q|`` and the distance.""" x_line = jnp.linspace(-L, L, n_points) x_plot = jnp.stack([x_line, jnp.zeros_like(x_line)], axis=1) fig, axes = plt.subplots(1, 3, figsize=(15, 5)) w_ex = jax.vmap(steady_state)(x_plot) z_plot = jax.vmap(default_bathymetry)(x_plot) q_mag_ex = jnp.sqrt(w_ex[:, 1] ** 2 + w_ex[:, 2] ** 2) variables.dofsl = dofsl_final w_h = variables.evaluate(x_plot) q_mag_h = jnp.sqrt(w_h[:, 1] ** 2 + w_h[:, 2] ** 2) axes[0].plot(x_line, w_ex[:, 0] + z_plot, "k--", label="steady h + Z") axes[0].plot(x_line, z_plot, "gray", linewidth=1.0, label="bathymetry Z") axes[0].plot(x_line, w_h[:, 0] + z_plot, label=f"enriched DG, t={T_FINAL}") axes[0].set_ylabel("h + Z") axes[0].set_title("Free surface") axes[1].plot(x_line, q_mag_ex, "k--", label="steady |Q|") axes[1].plot(x_line, q_mag_h, label=f"enriched DG, t={T_FINAL}") axes[1].set_ylabel("|Q|") axes[1].set_title("Discharge magnitude") axes[2].semilogy(x_line, jnp.abs(w_h[:, 0] - w_ex[:, 0]), label="|h_h - h|") axes[2].semilogy(x_line, jnp.abs(q_mag_h - q_mag_ex), label="||Q_h| - |Q||") axes[2].set_ylabel("distance to the steady state") axes[2].set_title( f"relative L2 at t={T_FINAL}: h {errors[0]:.2e}, Q {errors[1]:.2e}" ) for ax in axes: ax.set_xlabel("x1 (slice x2=0)") ax.legend(fontsize=7) ax.grid(True, alpha=0.3) fig.suptitle( f"2D shallow water, rotating steady vortex, {CATEGORY}\n" f"g={GRAVITY}, H0={H0}, alpha={ALPHA_DEFAULT}, gamma={GAMMA_DEFAULT}, " f"r0={R0_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, u)(x; alpha, gamma, r0) ...") _t0 = time.perf_counter() key, pinn = train_prior(key, file_name="pinn_prior_2d_shallow_water") 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}, {"velocity": 1}], solution=steady_h_and_u_wrt_mu, error=steady_h_and_u_wrt_mu, title="PINN prior on the steady (h, u) over the square, vs. closed form.", ) plt.show() plot_setup() prior_fn = make_prior_fn(pinn, MU_DEFAULT, postprocess=_swirl_postprocess) dt, nt = cfl_dt_nt(N_CELLS, ORDER, T_FINAL) mesh = cartesian_mesh( n_cells=list(N_CELLS), quad_order=ORDER + 3, bounds=BOX_BOUNDS ) variables = make_variables( mesh, ORDER, OUT_DIM, prior_fn, CATEGORY, multiplicative_indices=MULTIPLICATIVE_INDICES, ) scheme = make_scheme(variables, dt) dofsl_init, dofsl_final = solve( scheme, steady_state, nt, description=f"DG solve ({CATEGORY})" ) x_grid = error_grid() groups = ((0,), (1, 2)) # h, and Q as one vector err_init = relative_l2_errors(variables, dofsl_init, steady_state, x_grid, groups) err_final = relative_l2_errors(variables, dofsl_final, steady_state, x_grid, groups) 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_final, err_final)