r"""Solves the 2D compressible Euler equations with a discrete PINN. Test case: a smooth (C^1), Mach-dependent Gresho vortex, exact steady solution of the classical *multiscale* (low-Mach) Euler system .. math:: \partial_t \rho + \nabla\cdot(\rho u) &= 0, \\ \partial_t (\rho u) + \nabla\cdot(\rho u\otimes u) + \frac{1}{M^2}\nabla p &= 0, \\ \partial_t E + \nabla\cdot(u(E+p)) &= 0, with :math:`E = \rho e_{int} + M^2\,\tfrac12\rho|u|^2` and :math:`p = (\gamma-1)\rho e_{int}` (see :class:`SteadyEuler` in ``euler.py`` for the general M-parameterized residual this maps to). M carries *all* of the low-Mach stiffness here, explicitly, as a coefficient of the PDE — unlike the classical (Miczek) Gresho-vortex convention, which instead makes the background pressure huge (:math:`p_0=\rho/(\gamma\, \mathrm{Ma}^2)`) in the *unscaled* (M=1) Euler equations. Combining both (a huge Miczek background *inside* the M-parameterized momentum equation) double-counts the stiffness and drives the pressure fluctuation down to an absurd :math:`O(M^4)` *relative* to the background — this is why P0 here is the modest, M-independent :math:`\rho_0/\gamma`, not :math:`\rho_0/(\gamma M^2)`. With :math:`\xi = x - x_0`, :math:`\eta = y - y_0`, :math:`r = \sqrt{\xi^2+\eta^2}`, uniform density :math:`\rho = \rho_0`, the azimuthal velocity is the smooth, piecewise-cubic profile .. math:: u_\phi(r) = \begin{cases} 75\,r^2 - 250\,r^3 & 0 \le r \le 0.2 \\ -4 + 60\,r - 225\,r^2 + 250\,r^3 & 0.2 < r < 0.4 \\ 0 & r \ge 0.4 \end{cases}, \qquad (u, v) = \frac{u_\phi(r)}{r}\,(-\eta, \xi), and the pressure is :math:`p(r) = P_0 + M^2\,\mathrm{poly}(r)`, where ``poly`` (see ``exact_pressure``) satisfies :math:`d(\mathrm{poly})/dr = u_\phi(r)^2/r` exactly — i.e. the steady radial momentum balance :math:`-\rho u_\phi(r)^2/r + (1/M^2)\,dp/dr = 0` holds exactly, so this remains an exact **steady** solution of the full multiscale Euler system for any M. Continuity and the energy equation are satisfied automatically for *any* radial profile here, since any azimuthal field :math:`h(r)(-\eta,\xi)` is exactly divergence-free. The velocity and pressure fluctuation are exactly uniform beyond :math:`r = 0.4`, well inside the domain, so periodic boundary conditions are effectively exact. The solution is time-independent, so the error at any time :math:`t > 0` measures how well the discrete PINN preserves the steady state. Since M (not a huge background) carries the stiffness, the state :math:`W=(\rho,m_x,m_y,E)` stays O(1) everywhere by construction — none of ``post_processing_w`` / ``p_background`` / ``SOLVE_WEIGHTS`` below are load -bearing for numerical conditioning the way they were under the old Miczek-P0 convention; they are kept mostly for interface consistency and documentation of the (now largely dormant) mechanism. What *does* remain genuinely stiff is momentum's :math:`(1/M^2)(\gamma-1)\nabla E` term (irreducibly, since it is M itself, not a background level, doing the scaling) — and it should *not* be down-weighted in the loss: down -weighting it removes real dynamical signal (it is the actual pressure -gradient force balance), unlike the old huge ``h_bg`` term, which was mostly a redundant divergence penalty. Verified empirically. """ 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.nonlinear_approximation.approximation_spaces.approximation_spaces import ( ApproximationSpace, ) from scimba_jax.nonlinear_approximation.integration.monte_carlo import ( DomainSampler, TensorizedSampler, ) from scimba_jax.nonlinear_approximation.networks.mlp import MLP from scimba_jax.nonlinear_approximation.numerical_solvers.discrete_pinns import ( DiscretePINN, ) from scimba_jax.physical_models.temporal_pde.euler import ( SteadyEuler, SteadyExplicitMultiscaleEuler, SteadyImplicitMultiscaleEuler, ) from scimba_jax.plots.plots_nd import plot_abstract_approx_spaces from scimba_jax.time_discrete.butcher_tableau import ( build_imex_euler_tableau, build_implicit_euler_tableau, ) # ── problem parameters ──────────────────────────────────────────────────────── GAMMA = 1.4 MACH = 0.05 # reference Mach number: sets the acoustic/convective scale separation RHO0 = 1.0 # uniform density # Background pressure. NOTE: this is *not* Miczek's Mach-dependent # P0 = rho/(gamma*Ma**2) -- that convention encodes low-Mach stiffness via a # huge background pressure in the *unscaled* (M=1) Euler equations, which is # a different (and, combined with the M-parameterized PDE below, doubly # stiff / degenerate) convention than the one used here. With the explicit # multiscale system (1/M**2) grad(p) in SteadyEuler etc., M alone # carries the stiffness, and the natural background pressure is O(1): # P0 = rho0/gamma (see exact_pressure for the derivation). P0 = RHO0 / GAMMA E0 = P0 / (GAMMA - 1.0) # constant background total energy (see p_background) R1, R2 = 0.2, 0.4 # radii delimiting the vortex core / ring / exterior X0, Y0 = 0.5, 0.5 # vortex center L = 1.0 # domain side length # ── numerical parameters ───────────────────────────────────────────────────── N_COLLOC = 2500 N_EPOCHS_INIT = 2000 N_EPOCHS = 15 NT = 20 DOM_X = Square2D([(0.0, L), (0.0, L)], is_main_domain=True) DOM_T = (0.0, 0.2) SAMPLER = TensorizedSampler([DomainSampler(DOM_X)], model_type="x") # With P0 = rho0/gamma now O(1) (see P0 above), h_bg = P0*gamma/(gamma-1) is # O(1) too, so this weight is close to 1 and largely a no-op -- kept mostly # for robustness/documentation of the mechanism, not because it is doing # heavy lifting the way it did under the old (Miczek) P0 ~ 1/Ma**2 scaling. # The genuinely stiff O(1/M**2) term now lives in momentum, not energy, and # should *not* be down-weighted (down-weighting it removes real dynamical # signal rather than a redundant penalty -- verified empirically). SOLVE_WEIGHTS = {"interior": [1.0, 1.0, 1.0, 1.0 / P0**2]} # ── exact steady solution ───────────────────────────────────────────────────── def _u_phi(r: jnp.ndarray) -> jnp.ndarray: """Smooth (C^1) azimuthal velocity profile, in the same "natural units" as the classical piecewise-linear Gresho profile (max value 1 at r=0.2). """ in_core = r <= R1 in_ring = jnp.logical_and(r > R1, r < R2) u_core = 75.0 * r**2 - 250.0 * r**3 u_ring = -4.0 + 60.0 * r - 225.0 * r**2 + 250.0 * r**3 return jnp.where(in_core, u_core, jnp.where(in_ring, u_ring, 0.0)) def exact_pressure(t: jnp.ndarray, x: jnp.ndarray) -> jnp.ndarray: """Pressure for the stationary, Mach-dependent, smooth (C^1) Gresho vortex. The r-dependent part below (``poly_*``) satisfies d(poly)/dr = u_phi(r)**2 / r exactly, in the same natural (M-independent) units as u_phi. For the multiscale momentum equation dt(rho u) + div(...) + (1/M**2) grad(p) = 0, the steady radial balance -rho u_phi**2/r + (1/M**2) dp/dr = 0 requires dp/dr = M**2 rho u_phi**2/r, i.e. the fluctuation must be scaled by M**2 relative to ``poly`` (whereas P0, being spatially constant, needs no such scaling — its own role is set independently by the p_background convention, see euler.py). Args: t: time, shape (..., 1). Not used (solution is steady). x: spatial coordinates, shape (..., 2). Returns: Pressure field of shape (..., 1). """ xi = x[..., 0:1] - X0 # x - x_0 eta = x[..., 1:2] - Y0 # y - y_0 r2 = xi**2 + eta**2 r = jnp.sqrt(r2 + 1e-14) # regularized to keep the log/division branches safe in_core = r <= R1 in_ring = jnp.logical_and(r > R1, r < R2) poly_core = 1406.25 * r**4 - 7500.0 * r**5 + (10416.0 + 2.0 / 3.0) * r**6 poly_ring = ( 65.8843399322788 - 480.0 * r + 2700.0 * r**2 - (9666.0 + 2.0 / 3.0) * r**3 + 20156.25 * r**4 - 22500.0 * r**5 + (10416.0 + 2.0 / 3.0) * r**6 + 16.0 * jnp.log(r) ) poly_ext = ( 65.8843399322788 - 480.0 * R2 + 2700.0 * R2**2 - (9666.0 + 2.0 / 3.0) * R2**3 + 20156.25 * R2**4 - 22500.0 * R2**5 + (10416.0 + 2.0 / 3.0) * R2**6 + 16.0 * jnp.log(R2) ) poly = jnp.where(in_core, poly_core, jnp.where(in_ring, poly_ring, poly_ext)) return P0 + MACH**2 * poly def exact_sol(t: jnp.ndarray, x: jnp.ndarray) -> jnp.ndarray: """The stationary, Mach-dependent, smooth (C^1) Gresho vortex — exact for all t. Args: t: time, shape (..., 1). Not used (solution is steady). x: spatial coordinates, shape (..., 2). Returns: W = (rho, m_x, m_y, e) of shape (..., 4), where e = E - E0 is the energy fluctuation around the constant background E0 (see p_background in ``euler.py``). """ xi = x[..., 0:1] - X0 # x - x_0 eta = x[..., 1:2] - Y0 # y - y_0 r2 = xi**2 + eta**2 r = jnp.sqrt(r2 + 1e-14) # regularized to keep the log/division branches safe # u_phi(r) / r, so that (u, v) = (u_phi(r)/r) * (-eta, xi) is regular at r = 0 g = _u_phi(r) / r p = exact_pressure(t, x) rho = RHO0 * jnp.ones_like(r) u = -g * eta v = g * xi m_x = rho * u m_y = rho * v # kinetic energy scaled by M**2, consistent with the multiscale Euler # system's E = rho*e_int + M**2 * 0.5*rho*|u|**2 (see euler.py) E = p / (GAMMA - 1.0) + MACH**2 * 0.5 * rho * (u**2 + v**2) e_tilde = E - E0 return jnp.concatenate([rho, m_x, m_y, e_tilde], axis=-1) def f_init(x: jnp.ndarray) -> jnp.ndarray: """Initial condition W(t=0, x).""" return exact_sol(jnp.zeros_like(x[..., :1]), x) # ── network output reparametrization ───────────────────────────────────────── # P0 is now O(1) (see P0 above), so this is no longer about keeping a huge # background out of the network's raw output range -- but it is kept for # consistency with the p_background convention wired into SteadyEuler / # SteadyExplicitMultiscaleEuler / SteadyImplicitMultiscaleEuler # below (state's 4th component = energy fluctuation E - E0), and because it # is harmless. rho and momentum are reconstructed from perturbed primitive # variables (rho', u, v, p') as usual; the energy component uses the same # EOS as exact_sol / euler.py, with kinetic energy scaled by M**2. def post_processing_w(output: jnp.ndarray, x: jnp.ndarray) -> jnp.ndarray: rho = RHO0 + output[0] u = output[1] v = output[2] p_tilde = output[3] m_x = rho * u m_y = rho * v e_tilde = p_tilde / (GAMMA - 1.0) + MACH**2 * 0.5 * rho * (u**2 + v**2) return jnp.array([rho, m_x, m_y, e_tilde]) # ── PDEs ────────────────────────────────────────────────────────────────────── pde_full = SteadyEuler( main_domain=DOM_X, time_domain=DOM_T, gamma=GAMMA, p_background=P0, M=MACH, bc="strong", ) pde_exp = SteadyExplicitMultiscaleEuler( main_domain=DOM_X, time_domain=DOM_T, gamma=GAMMA, M=MACH, p_background=P0, bc="strong", ) pde_imp = SteadyImplicitMultiscaleEuler( main_domain=DOM_X, time_domain=DOM_T, gamma=GAMMA, M=MACH, p_background=P0, bc="strong", ) params = { "implicit_euler": { "tableau": build_implicit_euler_tableau(), "explicit_pde": None, "implicit_pde": pde_full, }, "implicit_explicit_euler": { "tableau": build_imex_euler_tableau(), "explicit_pde": pde_exp, "implicit_pde": pde_imp, }, } # ── main loop ───────────────────────────────────────────────────────────────── in_size = 2 # (x, y) out_size = 4 # (rho, m_x, m_y, e), e = E - E0 (see p_background) space = None spaces = [] losses = [] errors_over_time = [] for method, param in params.items(): butcher_tableau = param["tableau"] explicit_pde = param["explicit_pde"] implicit_pde = param["implicit_pde"] key = jax.random.PRNGKey(0) discrete_pinn = DiscretePINN( DOM_X, DOM_T, SAMPLER, out_size, NT, butcher_tableau, explicit_pde, implicit_pde, exact_solution=exact_sol, ) if space is None: nn = MLP( in_size=in_size, out_size=out_size, hidden_sizes=[16] * 3, key=key, embedding="periodic", periods=[L, L], n_periodic_features=3, ) space = ApproximationSpace( {"x": 2}, [(nn, "vec", out_size)], model_type="x", post_processing=post_processing_w, ) print("Initializing the discrete PINN...") key, discrete_pinn = discrete_pinn.initialize( key, space, f_init, N_EPOCHS_INIT, N_COLLOC, file_name="euler_gresho_vortex", retrain=False, ) space = discrete_pinn.space print("Initializing the discrete PINN... Done\n") plot_abstract_approx_spaces( [space], DOM_X, solution=f_init, error=f_init, loss=discrete_pinn.losses, title="Multiscale Gresho vortex — initial condition", ) plt.show() print(f"\nSolving with the {method} method...") key, discrete_pinn = discrete_pinn.solve( key, space, N_EPOCHS, N_COLLOC, weights=SOLVE_WEIGHTS ) spaces.append(discrete_pinn.space) losses.append(discrete_pinn.losses) errors_over_time.append(discrete_pinn.errors_over_time) print(f"\nSolving with the {method} method... Done\n") # ── error plots ─────────────────────────────────────────────────────────────── fig, ax = plt.subplots(1, 2, figsize=(12, 5)) error_type = ["L2", "Linf"] for i in range(2): for j, (method_name, _) in enumerate(params.items()): ax[i].semilogy( jnp.linspace(DOM_T[0], DOM_T[1], NT + 1), errors_over_time[j][:, i], label=method_name, ) ax[i].set_xlabel("time") ax[i].set_ylabel(f"relative {error_type[i]} error") ax[i].set_title(f"Relative {error_type[i]} error vs time") ax[i].legend() ax[i].grid() def compute_pressure_diff_exact(*vars): w = vars[0] rho, *_ = w.components() def p_diff_ex(x): return exact_pressure(None, x) - P0 # make it a ParamScalarFunction return p_diff_ex * rho / rho def compute_pressure_diff(*vars): w = vars[0] rho, q1, q2, e_tilde = w.components() norm_m2 = q1**2 + q2**2 e = e_tilde + E0 p = (GAMMA - 1) * (e - 0.5 * MACH**2 * norm_m2 / rho) return p - P0 def compute_pressure_error(*vars): dp = compute_pressure_diff(*vars) dp_exact = compute_pressure_diff_exact(*vars) exact_pressure = dp_exact + P0 return (dp - dp_exact) / exact_pressure additional_scalar_functions = { r"$p_{\text{exact}} - p_0$": compute_pressure_diff_exact, r"$p - p_0$": compute_pressure_diff, r"$p$ relative error": compute_pressure_error, } t_final = DOM_T[-1] plot_abstract_approx_spaces( spaces, DOM_X, loss=losses, title=f"Euler multiscale Gresho vortex at t = {t_final}", titles=[f"{method_name}" for method_name in params.keys()], additional_scalar_functions=additional_scalar_functions, ) plt.show()