r"""Solves the 2D isentropic incompressible Euler equations (vorticity-streamfunction formulation) for the double shear layer instability, using the Neural Semi-Lagrangian (NSL) predictor-corrector scheme. .. 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. At every macro time step :math:`[t^n, t^{n+1}]` the coupled system is solved with the predictor-corrector algorithm below, using :class:`NeuralSemiLagrangian` for both transport substeps and a standard elliptic PINN (:class:`Projector` + :class:`LaplacianDirichletND`) for the two Poisson solves: 1. solve :math:`-\Delta \psi^n = \omega^n` to get :math:`U^n = (\partial_y \psi^n, -\partial_x \psi^n)`; 2. solve the transport equation with :math:`U^n`, from :math:`\omega^n`, to get a prediction :math:`\omega^{*}`; 3. solve :math:`-\Delta \psi^{*} = \omega^{*}` to get :math:`U^{*}`, and set :math:`\bar{U} = (U^n + U^{*}) / 2`; 4. solve the transport equation with :math:`\bar{U}`, from :math:`\omega^n` (not :math:`\omega^{*}`), to get :math:`\omega^{n+1}`. The initial condition is the classical double shear layer test case (Bell, Colella & Glaz, 1989): two anti-parallel, hyperbolic-tangent shear layers at :math:`y = \pi/2` and :math:`y = 3\pi/2`, perturbed by a small sinusoidal cross-flow that seeds the Kelvin-Helmholtz roll-up of each layer: .. math:: U_1(x, y, 0) = \begin{cases} \tanh((y - \pi/2)/\rho) & y \le \pi \\ \tanh((3\pi/2 - y)/\rho) & y > \pi \end{cases}, \qquad U_2(x, y, 0) = \delta \sin(x), with :math:`\rho` the (common) shear layer width and :math:`\delta` the perturbation amplitude; :math:`\omega_0 = \partial_x U_2 - \partial_y U_1` is computed analytically from this and used directly as the fitting target for the initial condition (see ``omega_init``). The trained ``omega`` :class:`ApproximationSpace` is checkpointed to disk at every macro time step (rather than rendering a figure inline) -- see ``euler_2d_double_shear_layer_post_process.py``, which reloads every checkpoint, makes the diagnostic plots, saves one PNG snapshot per step, and prints the ``ffmpeg`` command to stitch them into a video. Run this script first, then that one. The macro predictor-corrector loop itself runs under a single :func:`jax.lax.scan`, not a Python ``for`` loop -- this is what actually fixes this script being atrociously slow, and it is not just a matter of Python-loop overhead: ``solve_psi_pinn`` rebuilds a fresh :class:`Projector` every macro step (so does :meth:`NeuralSemiLagrangian.solve`, internally), and :class:`Projector`'s ``__init__`` calls ``jax.jit`` on brand-new Python closures every time; on top of that, ``solve_psi_pinn``'s ``residual._construct_rhs(...)`` rewrap dynamically defines *and pytree-registers* a brand-new Python class on every call. Under a plain Python loop, both costs are paid fresh on every one of ``NT`` macro steps (a full XLA trace+compile, times ~1200 over a full run). ``jax.lax.scan`` traces the step body exactly once and reuses the compiled program for every iteration, so all of that Python-level setup happens once, not once per step -- exactly what the :mod:`vlasov_poisson_1d_1v_bump_on_tail` sibling example already relies on. Because checkpointing (writing to disk) cannot happen inside a traced ``scan`` body, the scanned step instead returns the whole trained ``omega`` space as a stacked output (alongside the loss history, like the Vlasov-Poisson example already does for its own diagnostics); the checkpoints are then saved from that stacked array, in an ordinary Python loop, once the scan itself has finished. Set the environment variable ``DSL_SMOKE=1`` to run a hugely reduced smoke test (tiny networks, few collocation points/epochs, a handful of time steps) that exercises the full pipeline without the cost of the real run. """ # %% import os import time import warnings import jax import jax.numpy as jnp import numpy as np from tqdm import tqdm 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.neural_semi_lagrangian import ( NeuralSemiLagrangian, ) from scimba_jax.nonlinear_approximation.numerical_solvers.projectors import Projector from scimba_jax.physical_models.elliptic_pde.laplacians import LaplacianDirichletND 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 omega ApproximationSpace checkpoint per macro time step (postfix "00000", # "00001", ...), saved with the generic ScimbaPytree.save/load facility (see # save_omega_checkpoint below) under the default ~/.scimba/scimba_jax/ location -- # same convention already used for "double_shear_layer_init" and # "double_shear_layer_poisson_init" further down. Reloaded by the companion # euler_2d_double_shear_layer_post_process.py to make every plot/video frame. OMEGA_CKPT_SCRIPTNAME = "double_shear_layer_omega" # Smoke-test switch: DSL_SMOKE=1 python double_shear_layer_2d.py -> tiny networks, # few collocation points/epochs, a handful of time steps. See the module docstring. _SMOKE = int(os.environ.get("DSL_SMOKE", "0")) # ───────────────────────────────────────────────────────────────────────────── # Physical parameters of the double shear layer test case # ───────────────────────────────────────────────────────────────────────────── 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, label_str="omega" ) DOMAIN_BOUNDS = ((X_MIN,), (X_MAX,)) # square domain: one period covers both axes SAMPLER_OMEGA = TensorizedSampler([DomainSampler(DOMAIN_XY)], model_type="x") SAMPLER_PSI = TensorizedSampler([DomainSampler(DOMAIN_XY)], model_type="x") # Time domain and macro time-stepping DOM_T = (0.0, 12.0) DT = 0.04 NT = int(round((DOM_T[1] - DOM_T[0]) / DT)) N_RK_SUBSTEPS = 1 # RK4 substeps per NSL transport solve # Network sizes HIDDEN_OMEGA = [30] * 6 HIDDEN_PSI = [20] * 6 # Training hyperparameters N_COLLOC_OMEGA = 128**2 # collocation points for the (x, y) physical space N_COLLOC_PSI = 96**2 # collocation points for the streamfunction N_EPOCHS_INIT = 250 # epochs to fit the initial condition N_EPOCHS_OMEGA = 100 # epochs per vorticity transport substep N_EPOCHS_PSI_INIT = 350 # epochs for the very first Poisson solve N_EPOCHS_PSI = 60 # epochs for subsequent (warm-started) Poisson solves # 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. omega develops fine-scale filamentation (roll-up of the two shear layers) # as t grows, so it is given more harmonics than the (smoother, Poisson-filtered) # streamfunction psi. X_PERIOD = X_MAX - X_MIN Y_PERIOD = Y_MAX - Y_MIN N_PERIODIC_FEATURES_OMEGA = 8 N_PERIODIC_FEATURES_PSI = 4 # Grid resolution used for snapshotting/diagnostics (not for training). N_GRID_PLOT = 256 if _SMOKE: N_COLLOC_OMEGA = 8**2 N_COLLOC_PSI = 8**2 N_EPOCHS_INIT = 5 N_EPOCHS_OMEGA = 3 N_EPOCHS_PSI_INIT = 5 N_EPOCHS_PSI = 3 HIDDEN_OMEGA = [8] * 2 HIDDEN_PSI = [8] * 2 N_PERIODIC_FEATURES_OMEGA = 2 N_PERIODIC_FEATURES_PSI = 2 N_GRID_PLOT = 32 DOM_T = (0.0, 0.06) NT = 3 def build_space_omega(key: jnp.ndarray) -> ApproximationSpace: """Build a fresh omega ApproximationSpace (periodic embedding in x and 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_omega = MLP( in_size=2, out_size=1, activation="sin", hidden_sizes=HIDDEN_OMEGA, key=key, embedding="periodic", periods=(X_PERIOD, Y_PERIOD), embedding_axes=(0, 1), n_periodic_features=N_PERIODIC_FEATURES_OMEGA, ) return ApproximationSpace({"x": 2}, [(nn_omega, "scalar", None)], model_type="x") def omega_init(xy: jnp.ndarray) -> jnp.ndarray: """Double shear layer initial vorticity, computed analytically from U1, U2. 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 _omega_on_grid( space_omega, x_grid: jnp.ndarray, y_grid: jnp.ndarray ) -> jnp.ndarray: """Evaluate omega on the tensor grid (x_grid, y_grid); shape (len(x_grid), len(y_grid)). Shared by every consumer that needs omega on a regular grid: the PINN Poisson RHS quadrature (compute_omega_average) and the plotting/snapshot code. """ variables_omega = space_omega.create_variables() omega_batched = variables_omega[0].vmap_on_physical_variables() xx, yy = jnp.meshgrid(x_grid, y_grid, indexing="ij") xy = jnp.stack([xx.ravel(), yy.ravel()], axis=-1) omega_vals = omega_batched(space_omega, xy) return omega_vals.reshape(x_grid.shape[0], y_grid.shape[0]) def compute_omega_average(space_omega, n_grid: int = 128) -> jnp.ndarray: """Compute the domain average of omega on a regular grid. The exact continuous solution has zero domain average at every time (the initial condition does, and pure transport by a divergence-free, periodic field conserves it); this is only an explicit safety net against small drift of the trained network's own average, mirroring how the Poisson RHS is kept exactly compatible in the analogous Vlasov-Poisson example. """ grid = jnp.linspace(X_MIN, X_MAX, n_grid, endpoint=False) return jnp.mean(_omega_on_grid(space_omega, grid, grid)) def make_omega_rhs(space_omega, omega_average: jnp.ndarray): """Build the (pointwise) RHS x -> omega(x) - of the Poisson residual. The sign matches LaplacianResidual's convention (LHS = -Delta psi), so LHS = RHS enforces -Delta psi = omega, as required. """ omega_fn = space_omega.create_variables()[0] def omega_rhs(xy: jnp.ndarray) -> jnp.ndarray: return omega_fn(space_omega, xy) - omega_average return omega_rhs def solve_psi_pinn( model_psi, key, space_psi, space_omega, n_epochs: int, n_colloc: int, file_name: str | None = None, retrain: bool = True, ): """One elliptic PINN solve of -Delta psi = omega, warm-started from space_psi. ``model_psi`` must be a :class:`LaplacianDirichletND` built once, outside of any ``jax.lax.scan``/Python loop rebuild: constructing a fresh :class:`AbstractPhysicalModel` triggers a concrete-only boundary computation on ``main_domain`` that is incompatible with jax's tracing. Instead, we reuse one pre-built model and just mutate its ``f_rhs``. """ omega_average = compute_omega_average(space_omega) omega_rhs = make_omega_rhs(space_omega, omega_average) residual = model_psi.physical_residuals[DOMAIN_XY.get_label()] # Re-wrap through `_construct_rhs` (what `__init__` does for the initial f_rhs) # rather than assigning the raw closure directly: a raw Python function is not # a valid pytree leaf, which only bites when `proj_psi.save` below actually # checkpoints this model -- see `AbstractFunctionalField`. residual.f_rhs = residual._construct_rhs(omega_rhs) proj_psi = Projector(model_psi, space_psi, SAMPLER_PSI, **ENG_KWARGS) if (file_name is not None) and (not retrain): n_epochs_load, proj_psi = proj_psi.load(file_name) if n_epochs_load > 0: print(f"Loaded pre-trained model from {file_name}.") return proj_psi.key, proj_psi.space key, proj_psi = proj_psi.project_scan(key, space_psi, n_epochs, n_colloc) if file_name is not None: proj_psi.save(file_name) return key, proj_psi.space def make_velocity_field_pinn(space_psi): """Build a pointwise x -> U(x) = (d_y psi(x), -d_x psi(x)) function. Pointwise (single point in, single point out): this is how :class:`Characteristic` calls the advection field (it is vmapped/batched from the outside, as an RHS of the NSL transport residual). ``curl`` is exactly (d_y ., -d_x .) -- see :meth:`ParamScalarFunction.curl`. """ variables_psi = space_psi.create_variables() curl_psi = variables_psi[0].curl("x") def velocity_field(xy: jnp.ndarray) -> jnp.ndarray: return curl_psi(space_psi, xy) return velocity_field def make_advection_field(space_psi): """Advection field U(x) for the vorticity characteristics, from one streamfunction.""" velocity_field = make_velocity_field_pinn(space_psi) def advection_field(t, xy): return velocity_field(xy) return advection_field def make_averaged_advection_field(space_psi_a, space_psi_b): """Advection field (U_a(x) + U_b(x)) / 2, for the corrector substep.""" field_a = make_velocity_field_pinn(space_psi_a) field_b = make_velocity_field_pinn(space_psi_b) def advection_field(t, xy): return 0.5 * (field_a(xy) + field_b(xy)) return advection_field def omega_checkpoint_postfix(frame_idx: int) -> str: """Zero-padded postfix identifying the omega checkpoint at macro step `frame_idx`.""" return f"{frame_idx:05d}" def save_omega_checkpoint(space_omega: ApproximationSpace, frame_idx: int) -> None: """Save space_omega's weights, to be reloaded by the post-processing script. Uses the generic ``ScimbaPytree.save`` facility (weights only, no training bookkeeping) rather than rendering a figure inline -- see the module docstring. """ space_omega.save(OMEGA_CKPT_SCRIPTNAME, postfix=omega_checkpoint_postfix(frame_idx)) # %% if __name__ == "__main__": """Solve the double shear layer instability with Neural Semi-Lagrangian.""" key = jax.random.PRNGKey(0) # vorticity omega(x, y): periodic embedding in both axes (the domain is square). key, subkey = jax.random.split(key) space_omega = build_space_omega(subkey) # `psi_state_n`/`psi_state_star` each carry an independent, warm-started # streamfunction ApproximationSpace, exactly like `phi_state_n`/`phi_state_star` # in the analogous Vlasov-Poisson example -- periodic embedding makes psi # exactly periodic by construction, so no boundary condition/anchor is needed. key, subkey = jax.random.split(key) nn_psi_n = MLP( in_size=2, out_size=1, hidden_sizes=HIDDEN_PSI, key=subkey, embedding="periodic", periods=(X_PERIOD, Y_PERIOD), embedding_axes=(0, 1), n_periodic_features=N_PERIODIC_FEATURES_PSI, ) psi_state_n = ApproximationSpace( {"x": 2}, [(nn_psi_n, "scalar", None)], model_type="x" ) key, subkey = jax.random.split(key) nn_psi_star = MLP( in_size=2, out_size=1, hidden_sizes=HIDDEN_PSI, key=subkey, embedding="periodic", periods=(X_PERIOD, Y_PERIOD), embedding_axes=(0, 1), n_periodic_features=N_PERIODIC_FEATURES_PSI, ) psi_state_star = ApproximationSpace( {"x": 2}, [(nn_psi_star, "scalar", None)], model_type="x" ) # Built once, outside the main loop -- see solve_psi_pinn's docstring. model_psi = LaplacianDirichletND(DOMAIN_XY, bc="strong", model_type="x") def _placeholder_field(t, xy): return jnp.zeros_like(xy) nsl = NeuralSemiLagrangian( main_domain=DOMAIN_XY, time_domain=(0.0, DT), sampler=SAMPLER_OMEGA, dt=DT, advection_field=_placeholder_field, periodic=True, domain_bounds=DOMAIN_BOUNDS, out_size=1, model_type="x", n_rk_steps=N_RK_SUBSTEPS, ) print("Initializing omega^0 (double shear layer initial condition)...") start_init = time.perf_counter() key, nsl = nsl.initialize( key, space_omega, omega_init, N_EPOCHS_INIT, N_COLLOC_OMEGA, file_name="double_shear_layer_init", retrain=False, **ENG_KWARGS, ) end_init = time.perf_counter() space_omega = nsl.space print(f"Initializing omega^0... Done in {end_init - start_init:.2f} seconds\n") # cold-start Poisson solve for U^0 (untrained streamfunction net -> more epochs); # every later U^n solve, below, is warm-started and only needs N_EPOCHS_PSI. print("Solving the initial Poisson equation for psi^0...") start_psi0 = time.perf_counter() key, psi_state_n = solve_psi_pinn( model_psi, key, psi_state_n, space_omega, N_EPOCHS_PSI_INIT, N_COLLOC_PSI, file_name="double_shear_layer_poisson_init", retrain=False, ) end_psi0 = time.perf_counter() print(f"Solving for psi^0... Done in {end_psi0 - start_psi0:.2f} seconds\n") # psi_state_star is seeded with psi_state_n's freshly solved value -- cheap and # valid, since the very first predictor substep uses U^n (psi_state_n), not yet # psi_state_star. psi_state_star = psi_state_n def time_step(carry, _): """One predictor-corrector macro time step [t_n, t_n + DT]. Traced exactly once by the jax.lax.scan below (not once per macro step, unlike the Python for-loop this replaces) -- see the module docstring for why that is what actually fixes this script's performance: solve_psi_pinn and NeuralSemiLagrangian.solve each rebuild a Projector (hence a fresh jax.jit) every macro step, and solve_psi_pinn's _construct_rhs rewrap registers a brand-new pytree class every call. All of that expensive Python-level setup now happens once, at trace time, instead of NT times. """ key, space_omega, psi_state_n, psi_state_star = carry # 2. transport omega^n with U^n to get the prediction omega^* nsl.characteristic.advection_field = make_advection_field(psi_state_n) key, nsl_star = nsl.solve( key, space_omega, N_EPOCHS_OMEGA, N_COLLOC_OMEGA, no_tqdm=True, **ENG_KWARGS ) space_omega_star = nsl_star.space predictor_losses = nsl_star.losses.losses_history["total"][-N_EPOCHS_OMEGA:] # 3. solve -Delta psi^* = omega^* to get U^*, and set Ubar = (U^n + U^*) / 2 key, psi_state_star = solve_psi_pinn( model_psi, key, psi_state_star, space_omega_star, N_EPOCHS_PSI, N_COLLOC_PSI ) # 4. transport omega^n (not omega^*) with Ubar to get omega^{n+1} nsl.characteristic.advection_field = make_averaged_advection_field( psi_state_n, psi_state_star ) key, nsl_next = nsl.solve( key, space_omega, N_EPOCHS_OMEGA, N_COLLOC_OMEGA, no_tqdm=True, **ENG_KWARGS ) space_omega = nsl_next.space corrector_losses = nsl_next.losses.losses_history["total"][-N_EPOCHS_OMEGA:] # 1. (for the next iteration) solve -Delta psi^{n+1} = omega^{n+1}, # warm-started from psi^n key, psi_state_n = solve_psi_pinn( model_psi, key, psi_state_n, space_omega, N_EPOCHS_PSI, N_COLLOC_PSI ) new_carry = (key, space_omega, psi_state_n, psi_state_star) step_losses = jnp.concatenate([predictor_losses, corrector_losses], axis=0) # space_omega is stacked (not saved here): orbax checkpointing is not # traceable, so every step's trained space is instead returned as a scan # output and saved from the stacked array below, once the scan is done. return new_carry, (space_omega, step_losses) try: from jax_tqdm import scan_tqdm time_step = scan_tqdm(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." ) print(f"Running {NT} predictor-corrector steps...") start_loop = time.perf_counter() space_omega_0 = space_omega init_carry = (key, space_omega, psi_state_n, psi_state_star) ( (key, space_omega, psi_state_n, psi_state_star), (space_omega_history, transport_losses_history), ) = jax.lax.scan(time_step, init_carry, jnp.arange(NT)) end_loop = time.perf_counter() print(f"Running {NT} steps... Done in {end_loop - start_loop:.2f} seconds\n") losses_history = jnp.concatenate( [nsl.losses.losses_history["total"], transport_losses_history.reshape((-1,))], axis=0, ) np.save(LOSSES_HISTORY_FILE, np.asarray(losses_history)) print(f"\nSaving {NT + 1} omega checkpoints...") save_omega_checkpoint(space_omega_0, 0) for step in tqdm(range(NT)): save_omega_checkpoint( jax.tree_util.tree_map(lambda leaf: leaf[step], space_omega_history), step + 1, ) print( f"\n{NT + 1} omega checkpoints saved (scriptname={OMEGA_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 omega snapshots, and get the ffmpeg command." ) # %%