r"""Solves the 2D isentropic incompressible Euler equations (vorticity-streamfunction formulation) for the double shear layer instability, using a *discrete PINN* that solves for :math:`\omega` and :math:`\psi` **simultaneously**, at every implicit time step -- unlike the sibling ``neural_sl/advanced_examples/euler_2d_double_shear_layer.py``, which splits every macro step into separate Neural-Semi-Lagrangian transport substeps and elliptic Poisson solves. .. math:: \partial_t \omega(t, x, y) + U(t, x, y) \cdot \nabla \omega(t, x, y) = 0, \quad \forall (t, x, y) \in [0, T] \times \Omega, -\Delta \psi(t, x, y) = \omega(t, x, y), \quad U(t, x, y) = (\partial_y \psi(t, x, y), -\partial_x \psi(t, x, y)), with :math:`\Omega = (0, 2\pi)^2` periodic in both directions, :math:`\omega` the vorticity, :math:`\psi` the streamfunction, and :math:`U = (U_1, U_2)` the (divergence-free, by construction) velocity field. Same physics, initial condition (Bell, Colella & Glaz, 1989) and physical constants (:math:`\rho`, :math:`\delta`) as the NSL sibling -- see its module docstring for the physical discussion; only the numerical method differs. Why this is a DAE, not an ODE ------------------------------ :math:`\psi` carries **no time derivative** -- it is a pure algebraic (elliptic) constraint on :math:`\omega`, coupled to the transport equation. A standard discrete PINN (:class:`~scimba_jax.nonlinear_approximation.numerical_solvers. discrete_pinns.DiscretePINN`) gives every output component an *implicit identity mass matrix* (``rho + a_ii*dt*L(rho)``, see ``build_pdes``/ ``build_rhs_general_rk`` in ``discrete_pinns.py``), so treating :math:`(\omega, \psi)` as a plain 2-vector state would silently turn the Poisson equation into the *relaxed* :math:`\psi + a_{ii} dt (-\Delta\psi - \omega) = \psi^n`, which only recovers the true constraint as :math:`dt \to 0`. :class:`VorticityStreamfunctionDiscretePINN` below fixes this with a **diagonal 0/1 mass matrix**, ``mass = (1, 0)``: :math:`\omega`'s row keeps the usual identity-mass implicit treatment, :math:`\psi`'s row drops it entirely, so its stage equation reduces to :math:`a_{ii} dt (-\Delta\psi - \omega) = 0`, i.e. :math:`-\Delta \psi = \omega` -- independent of :math:`dt`, for *any* stage with :math:`a_{ii} \neq 0`. (A stage with :math:`a_{ii} = 0`, e.g. Crank- Nicolson's first stage as written in ``butcher_tableau.py``, would leave that stage's :math:`\psi` unconstrained -- avoid such tableaus here.) This only overrides ``DiscretePINN.build_pdes``/``build_rhs_general_rk`` (the two spots that hardcode the identity mass), kept local to this example since no other example needs a non-identity mass matrix. One direct benefit: the "carry" term ``u_n`` in ``build_rhs_general_rk`` is also masked out for :math:`\psi`, so its value at the *previous* stage never leaks into the new stage equation -- the Poisson constraint is refit from scratch, from the *current* stage's :math:`\omega`, every time. That also means :math:`\psi`'s initial condition (``state_init`` below) is never used. A second wrinkle: with ``mass=0``, :math:`\psi`'s *entire* stage equation is :math:`a_{ii} dt (-\Delta\psi - \omega)` -- :math:`O(dt)`, unlike :math:`\omega`'s row, which keeps an :math:`O(1)` anchor from its own identity term. Left alone, the least-squares fit underweights the Poisson constraint by :math:`O(dt^2)` relative to the transport equation. Rather than compensating with a ``weights=`` argument tuned to :math:`1/dt^2` (fragile: it would need to track ``DT`` by hand), ``dt`` is baked directly into :math:`\psi`'s own residual as :math:`(-\Delta\psi - \omega) / dt` (see ``VorticityStreamfunctionResidual``): the two ``dt`` factors cancel exactly, algebraically, at every stage of *any* implicit RK scheme with :math:`a_{ii} \neq 0` -- leaving :math:`-\Delta\psi - \omega` at the same :math:`O(1)` scale as the transport row, not an approximation, and no tuning knob to keep in sync. Time stepping: Alexander (3-stage, 3rd order, L-stable) ----------------------------------------------------- ``build_alexander_tableau()`` solves the whole coupled system at each of its three implicit stages -- what makes solving both fields *at once* possible in the first place. ``DT``/``NT``/the physical time horizon are kept numerically identical to the NSL sibling for a direct comparison of the *coupling* strategy (monolithic DAE vs. operator splitting). Like the NSL sibling, the trained state ``ApproximationSpace`` is checkpointed to disk once per macro step (see ``euler_2d_double_shear_layer_post_process.py``) rather than plotted inline, and the whole ``NT``-step loop runs under one ``jax.lax.scan`` (:meth:`VorticityStreamfunctionDiscretePINN.solve`, overridden from :class:`DiscretePINN` only to additionally stack the per-step trained space as a scan output, exactly like the NSL sibling's own hand-rolled scan) so that the one-time cost of building the ``Projector``/``jax.jit`` machinery is paid once, not ``NT`` times. Set the environment variable ``DSL_SMOKE=1`` to run a hugely reduced smoke test (tiny network, few collocation points/epochs, a handful of time steps). """ # %% import copy import os import warnings from typing import Callable import jax import jax.numpy as jnp import numpy as np from tqdm import tqdm from scimba_jax.domains.meshless_domains.base import VolumetricDomain from scimba_jax.domains.meshless_domains.domains_2d import Square2D from scimba_jax.nonlinear_approximation.approximation_spaces.approximation_spaces import ( AbstractApproxSpace, ApproximationSpace, ) from scimba_jax.nonlinear_approximation.integration.monte_carlo import ( DomainSampler, TensorizedSampler, ) from scimba_jax.nonlinear_approximation.model_class.funcparam_vectorial import ( ParamVecFunction, ) from scimba_jax.nonlinear_approximation.networks.mlp import MLP from scimba_jax.nonlinear_approximation.numerical_solvers.discrete_pinns import ( DiscretePINN, ) from scimba_jax.nonlinear_approximation.numerical_solvers.evaluators import Evaluator from scimba_jax.nonlinear_approximation.numerical_solvers.projectors import ( KEY_TYPE, Projector, ) from scimba_jax.nonlinear_approximation.optimizers.losses import ProjectorLosses from scimba_jax.physical_models.abstract_physical_model import AbstractPhysicalModel from scimba_jax.physical_models.abstract_residuals import ( PARAM_FUNC_TYPE, InteriorResidual, ) from scimba_jax.time_discrete.butcher_tableau import build_alexander_tableau DATA_DIR = os.path.join( os.path.dirname(os.path.abspath(__file__)), "data", "double_shear_layer" ) os.makedirs(DATA_DIR, exist_ok=True) LOSSES_HISTORY_FILE = os.path.join(DATA_DIR, "losses_history.npy") # One state (omega, psi) ApproximationSpace checkpoint per macro time step (postfix # "00000", "00001", ...), saved with the generic ScimbaPytree.save/load facility -- # reloaded by the companion euler_2d_double_shear_layer_post_process.py. STATE_CKPT_SCRIPTNAME = "double_shear_layer_discrete_pinn_state" # Smoke-test switch: DSL_SMOKE=1 python euler_2d_double_shear_layer.py -> tiny # network, few collocation points/epochs, a handful of time steps. _SMOKE = int(os.environ.get("DSL_SMOKE", "0")) # ───────────────────────────────────────────────────────────────────────────── # Physical parameters of the double shear layer test case -- identical to the # NSL sibling (neural_sl/advanced_examples/euler_2d_double_shear_layer.py). # ───────────────────────────────────────────────────────────────────────────── X_MIN, X_MAX = 0.0, 2 * jnp.pi Y_MIN, Y_MAX = 0.0, 2 * jnp.pi # Shear layer width and perturbation amplitude. RHO = jnp.pi / 15 DELTA = 0.05 DOMAIN_XY = Square2D([(X_MIN, X_MAX), (Y_MIN, Y_MAX)], is_main_domain=True) SAMPLER = TensorizedSampler([DomainSampler(DOMAIN_XY)], model_type="x") # Time domain and (single-stage, fully implicit) macro time-stepping. Kept # numerically identical to the NSL sibling -- see the module docstring's # trade-off discussion. DOM_T = (0.0, 12.0) DT = 0.04 NT = int(round((DOM_T[1] - DOM_T[0]) / DT)) # Network size: a single MLP produces both omega and psi (out_size=2), so its # capacity is a compromise between the NSL sibling's omega network (more # harmonics, for the filamentation of the roll-up) and its psi network # (smoother, Poisson-filtered). HIDDEN_STATE = [30] * 6 N_PERIODIC_FEATURES = 6 # Training hyperparameters N_COLLOC = 128**2 # collocation points, shared by both residual rows N_EPOCHS_INIT = 250 # epochs to fit the initial condition N_EPOCHS = 100 # epochs per implicit stage (2 stages per macro time step) # Natural gradient preconditioning parameters LINESEARCH = "strong-wolfe" MATRIX_REGULARIZATION = 1e-2 ADAPTIVE_MATRIX_REGULARIZATION = True ENG_KWARGS = { "linesearch": LINESEARCH, "matrix_regularization": MATRIX_REGULARIZATION, "adaptive_matrix_regularization": ADAPTIVE_MATRIX_REGULARIZATION, } # Periodic embedding in x AND y: the domain is square, so both axes share the # same period. X_PERIOD = X_MAX - X_MIN Y_PERIOD = Y_MAX - Y_MIN # Grid resolution used for snapshotting/diagnostics (not for training). N_GRID_PLOT = 256 if _SMOKE: N_COLLOC = 8**2 N_EPOCHS_INIT = 5 N_EPOCHS = 3 HIDDEN_STATE = [8] * 2 N_PERIODIC_FEATURES = 2 N_GRID_PLOT = 32 DOM_T = (0.0, 0.06) NT = 3 def build_space_state(key: jnp.ndarray) -> ApproximationSpace: """Build a fresh (omega, psi) ApproximationSpace (periodic embedding in x, y). Shared by the training loop and the post-processing script, so both build the exact same architecture -- required for a checkpoint saved by one to load into the other (only the weights differ; ``key`` only seeds the untrained case). """ nn_state = MLP( in_size=2, out_size=2, activation="sin", hidden_sizes=HIDDEN_STATE, key=key, embedding="periodic", periods=(X_PERIOD, Y_PERIOD), embedding_axes=(0, 1), n_periodic_features=N_PERIODIC_FEATURES, ) return ApproximationSpace({"x": 2}, [(nn_state, "vec", 2)], model_type="x") def omega_init(xy: jnp.ndarray) -> jnp.ndarray: """Double shear layer initial vorticity -- identical to the NSL sibling. omega_0 = d(U2)/dx - d(U1)/dy, with U1/U2 as given in the module docstring. Args: xy: A single (x, y) point of shape (2,). Returns: Initial vorticity value of shape (1,). """ x, y = xy[0:1], xy[1:2] arg_lower = (y - jnp.pi / 2) / RHO arg_upper = (3 * jnp.pi / 2 - y) / RHO du1_dy_lower = (1 - jnp.tanh(arg_lower) ** 2) / RHO du1_dy_upper = -(1 - jnp.tanh(arg_upper) ** 2) / RHO du1_dy = jnp.where(y <= jnp.pi, du1_dy_lower, du1_dy_upper) du2_dx = DELTA * jnp.cos(x) return du2_dx - du1_dy def state_init(xy: jnp.ndarray) -> jnp.ndarray: """Initial (omega, psi): psi's value is a placeholder, never actually used. Because psi's row has mass=0 (see module docstring), the "carry" term u_n is masked out of psi's stage equation at every step, including the first: psi^0 never enters the fit. Only omega's initial condition is physical. """ omega0 = omega_init(xy) psi0 = jnp.zeros_like(omega0) return jnp.concatenate([omega0, psi0], axis=-1) def state_checkpoint_postfix(frame_idx: int) -> str: """Zero-padded postfix identifying the state checkpoint at macro step `frame_idx`.""" return f"{frame_idx:05d}" def save_state_checkpoint(space_state: ApproximationSpace, frame_idx: int) -> None: """Save space_state's weights, to be reloaded by the post-processing script.""" space_state.save(STATE_CKPT_SCRIPTNAME, postfix=state_checkpoint_postfix(frame_idx)) # ───────────────────────────────────────────────────────────────────────────── # The (omega, psi) DAE residual, and a DiscretePINN that knows how to solve it. # ───────────────────────────────────────────────────────────────────────────── class VorticityStreamfunctionResidual(InteriorResidual): r"""Interior residual L(omega, psi), packed as a 2-vector -- see module docstring. ``construct_residual`` returns:: L_omega = U . grad(omega), U = curl(psi) = (d_y psi, -d_x psi) L_psi = (-Laplacian(psi) - omega) / dt so that, wrapped by :class:`VorticityStreamfunctionDiscretePINN`'s mass=(1, 0) backward-Euler stage equation ``mass*rho + dt*L(rho) = mass*u_n``, the omega row is the usual ``omega + dt*L_omega = omega^n``, i.e. the backward-Euler discretization of ``d(omega)/dt + U . grad(omega) = 0`` (transport), and the psi row is exactly ``dt * L_psi = 0``, i.e. (the ``/dt`` in L_psi and the ``dt`` from the stage equation cancelling algebraically) exactly ``-Laplacian(psi) = omega``, independent of dt. Args: domain: The spatial domain of the problem. time_domain: The time domain of the problem. dt: The (fixed) macro time step -- see class docstring for why this must be given to the residual itself, not just to the outer solver. model_type: The type of the model. Defaults to "x" (no explicit time input: this is a semi-discrete, discrete-PINN residual). """ def __init__( self, domain: VolumetricDomain, time_domain: tuple[float, float], dt: float, model_type: str = "x", ): super().__init__( domain=domain, size=2, model_type=model_type, f_rhs=None, time_domain=time_domain, ) self.dt = dt def construct_residual(self, *vars: PARAM_FUNC_TYPE) -> PARAM_FUNC_TYPE: state = vars[0] omega, psi = state.components() u = psi.curl("x") # (d_y psi, -d_x psi) l_omega = omega.advection("x", u) # U . grad(omega) l_psi = (-psi.laplacian("x") - omega) / self.dt return ParamVecFunction.cat((l_omega, l_psi)) class VorticityStreamfunctionEuler(AbstractPhysicalModel): """Physical model wrapping :class:`VorticityStreamfunctionResidual`. Periodic in both directions (see ``build_space_state``'s periodic MLP embedding): no boundary residual is added, so `bc` is always "strong" here -- there is no weak option, unlike e.g. :class:`SteadyHeatND`. Args: main_domain: The spatial domain of the problem (must be periodic-ready, i.e. its boundary conditions handled entirely by the network). time_domain: The time domain of the problem. dt: The macro time step, forwarded to :class:`VorticityStreamfunctionResidual`. model_type: The type of the model. Defaults to "x". """ def __init__( self, main_domain: VolumetricDomain, time_domain: tuple[float, float], dt: float, model_type: str = "x", ): super().__init__(main_domain=main_domain, time_domain=time_domain) self.physical_residuals = { self.main_domain.get_label(): VorticityStreamfunctionResidual( domain=main_domain, time_domain=time_domain, dt=dt, model_type=model_type, ) } class VorticityStreamfunctionDiscretePINN(DiscretePINN): """DiscretePINN with a diagonal 0/1 mass matrix -- see module docstring. Overrides the two spots in :class:`DiscretePINN` (``build_pdes``, ``build_rhs_general_rk``) that hardcode an identity mass matrix on every output component, plus ``solve`` (to additionally stack the per-step trained space as a scan output, for checkpointing -- exactly mirroring what the NSL sibling's own hand-rolled ``jax.lax.scan`` already does). Kept local to this example rather than touching ``discrete_pinns.py``: no other example in the library needs a non-identity mass matrix. Args: *args, **kwargs: forwarded to :class:`DiscretePINN`. mass: 1D array of shape ``(out_size,)``, one entry per output component: ``1.0`` for a genuinely time-evolving component, ``0.0`` for a pure algebraic constraint. """ def __init__(self, *args, mass: jnp.ndarray, **kwargs): self.mass = mass super().__init__(*args, **kwargs) def build_pdes(self): """Like :meth:`DiscretePINN.build_pdes`, but masking the identity term. The only change from the base implementation is ``rho * self.mass`` instead of ``rho`` -- see class/module docstrings. """ self.pdes = [] for current_stage in range(1, self.n_stages + 1): pde = copy.deepcopy(self.implicit_pde) old_construct_residual = self.get_pde_residual(pde) def new_construct_residual( *vars: PARAM_FUNC_TYPE, _stage: int = current_stage, _old_construct_residual=old_construct_residual, ) -> PARAM_FUNC_TYPE: rho = vars[0] a_imp_ii = self.butcher_tableau.compute_a_imp_ii(_stage) return rho * self.mass + a_imp_ii * self.dt * _old_construct_residual( *vars ) self.set_pde_residual(pde, new_construct_residual) self.pdes.append(pde) def build_rhs_general_rk( self, spaces: list[AbstractApproxSpace], t: float, stage: int ) -> Callable: """Like :meth:`DiscretePINN.build_rhs_general_rk`, but masking ``u_n``. The only change from the base implementation is ``self.mass * u_n`` instead of ``u_n`` in ``rhs_func`` -- see class/module docstrings. Every other term (``sum_residuals``, ``rhs``) is already a genuine residual/ source evaluation, not a carried-over state value, so it needs no mask. """ pde_exp, pde_imp = self.explicit_pde, self.implicit_pde explicit_rhs, implicit_rhs = self.explicit_rhs, self.implicit_rhs butcher_tableau = self.butcher_tableau dt = self.dt c_exp, c_imp = butcher_tableau.c_exp, butcher_tableau.c_imp def build_interior_residual(pde, var): res = pde.physical_residuals["interior"] return res.pre_processed_construct_residual(*var) def eval_rhs(rhs, c, *args): return rhs(t + c * dt, *args) if rhs is not None else 0 stages_data_exp = [] stages_data_imp = [] for space in spaces: var = space.create_variables() ev = Evaluator.construct_pointwise_evaluation(var) res_exp = build_interior_residual(pde_exp, var) res_imp = build_interior_residual(pde_imp, var) stages_data_exp.append({"variable": ev, "residual": res_exp}) stages_data_imp.append({"residual": res_imp}) exp_terms, imp_terms = butcher_tableau.stage_terms(stage) def compute_update(*args): sum_residuals = 0 for w, j, c_idx in exp_terms: res_j = stages_data_exp[j]["residual"](spaces[j], *args) rhs_j = eval_rhs(explicit_rhs, c_exp[c_idx], *args) sum_residuals += w * (res_j - rhs_j) for w, j, c_idx in imp_terms: res_j = stages_data_imp[j]["residual"](spaces[j], *args) rhs_j = eval_rhs(implicit_rhs, c_imp[c_idx], *args) sum_residuals += w * (res_j - rhs_j) return sum_residuals def rhs_func(*args): u_n = self.mass * stages_data_exp[0]["variable"](spaces[0], *args) sum_residuals = compute_update(*args) a_ii = butcher_tableau.compute_a_imp_ii(stage) rhs = eval_rhs(implicit_rhs, c_imp[stage - 1], *args) return u_n - dt * sum_residuals + a_ii * dt * rhs return rhs_func def solve( self, key: jnp.ndarray, space: AbstractApproxSpace, n_epochs: int, n_colloc: int, **kwargs, ) -> tuple[KEY_TYPE, "VorticityStreamfunctionDiscretePINN"]: """Like :meth:`DiscretePINN.solve`, but also stacking the per-step space. The only change from the base implementation is the extra ``spaces[-1]`` scan output (``space_history``): checkpointing to disk cannot happen inside the traced ``jax.lax.scan`` body, so (exactly like the NSL sibling) the whole trained space is instead returned as a stacked pytree, and the per-step checkpoints are saved from it in an ordinary Python loop once the scan has finished -- see the ``__main__`` block below. """ times = jnp.arange(self.nt) * self.dt + self.time_domain[0] key, l2_init, linf_init = self.compute_relative_error(key, space, times[0]) def time_step(carry, x): key, space = carry step_index, current_time = x spaces = [space] stage_losses = [] for pde, rhs_builder in zip(self.pdes, self.rhs_builders): self.set_pde_rhs(pde, rhs_builder(spaces, current_time)) proj = Projector(pde, spaces[-1], self.sampler, **kwargs) key, proj = proj.project_scan( key, spaces[-1], n_epochs, n_colloc, **kwargs ) spaces.append(proj.space) stage_losses.append(proj.losses.losses_history["total"]) key, l2, linf = self.compute_relative_error( key, spaces[-1], current_time + self.dt ) step_losses = jnp.concatenate(stage_losses, axis=0) return (key, spaces[-1]), (jnp.array([l2, linf]), step_losses, spaces[-1]) try: from jax_tqdm import scan_tqdm time_step = scan_tqdm(self.nt, print_rate=1, tqdm_type="std")(time_step) except ImportError: warnings.warn( "jax_tqdm not installed, cannot display tqdm progress bar. To " "display a progress bar, install jax_tqdm." ) (key, space), (errors, step_losses, space_history) = jax.lax.scan( time_step, (key, space), (jnp.arange(self.nt), times) ) initial_errors = jnp.array([[l2_init, linf_init]]) solve_losses_history = step_losses.reshape((-1,)) base_losses = self.losses if self.losses is not None else ProjectorLosses({}) init_losses_history = base_losses.losses_history.get( "total", jnp.array([], dtype=solve_losses_history.dtype) ) full_losses_history = jnp.concatenate( [init_losses_history, solve_losses_history], axis=0 ) new_discrete_pinn = self._with_space(space) new_discrete_pinn.errors_over_time = jnp.vstack([initial_errors, errors]) new_discrete_pinn.losses = base_losses.set_losses_history( {"total": full_losses_history} ) new_discrete_pinn.space_history = space_history return key, new_discrete_pinn # %% if __name__ == "__main__": """Solve the double shear layer instability with a discrete PINN DAE.""" key = jax.random.PRNGKey(0) key, subkey = jax.random.split(key) space_state = build_space_state(subkey) pde = VorticityStreamfunctionEuler(DOMAIN_XY, DOM_T, dt=DT, model_type="x") discrete_pinn = VorticityStreamfunctionDiscretePINN( DOMAIN_XY, DOM_T, SAMPLER, out_size=2, nt=NT, implicit_pde=pde, mass=jnp.array([1.0, 0.0]), butcher_tableau=build_alexander_tableau(), ) print("Initializing (omega^0, psi^0)...") key, discrete_pinn = discrete_pinn.initialize( key, space_state, state_init, N_EPOCHS_INIT, N_COLLOC, file_name="double_shear_layer_discrete_pinn_init", retrain=False, **ENG_KWARGS, ) space_state_0 = discrete_pinn.space print("Initializing (omega^0, psi^0)... Done\n") print(f"Running {NT} fully implicit (omega, psi) macro steps...") key, discrete_pinn = discrete_pinn.solve( key, space_state_0, N_EPOCHS, N_COLLOC, **ENG_KWARGS ) print(f"Running {NT} steps... Done\n") losses_history = discrete_pinn.losses.losses_history["total"] np.save(LOSSES_HISTORY_FILE, np.asarray(losses_history)) print(f"\nSaving {NT + 1} state checkpoints...") save_state_checkpoint(space_state_0, 0) for step in tqdm(range(NT)): save_state_checkpoint( jax.tree_util.tree_map( lambda leaf: leaf[step], discrete_pinn.space_history ), step + 1, ) print( f"\n{NT + 1} state checkpoints saved (scriptname={STATE_CKPT_SCRIPTNAME!r}, " f"postfix 00000..{NT:05d}) under ~/.scimba/scimba_jax/" ) print(f"Training loss history saved to {LOSSES_HISTORY_FILE}") print( "Run euler_2d_double_shear_layer_post_process.py to make the plots, save " "the snapshots, and get the ffmpeg command." ) # %%