r"""Post-processing companion to ``euler_2d_double_shear_layer.py``. That script trains the double shear layer discrete-PINN DAE solve and, at every macro time step, checkpoints the trained (omega, psi) ``ApproximationSpace`` to disk (via the generic ``ScimbaPytree.save`` facility) instead of rendering a figure inline. This script does no training: it rebuilds an (untrained) state space with the exact same architecture, reloads every saved checkpoint into it, and from each one * saves one PNG snapshot of omega (``fig/double_shear_layer/omega_%05d.png``), * recomputes the enstrophy diagnostic (omega only), * recomputes the DAE constraint violation ``max |Delta psi + omega|`` -- see below, then plots the enstrophy history, the constraint-violation history, and the training loss history saved alongside the checkpoints, and prints the ``ffmpeg`` command to stitch the PNG snapshots into a video. Run ``euler_2d_double_shear_layer.py`` first. Why the constraint-violation diagnostic ----------------------------------------- Unlike the NSL sibling (which solves a separate Poisson PINN for psi at every macro step, and so is only ever as consistent with omega as that solve's own training loss), this discrete-PINN solve fits omega and psi *jointly*: there is no dedicated Poisson solve to inspect in isolation. ``max |Delta psi + omega|`` on a regular grid is the direct measure of how well the network satisfies the algebraic constraint at each checkpoint -- the quantity the whole point of this example (solving both fields at once, see the training script's module docstring) is trying to keep small without a separate solve driving it down. """ # %% import glob import os import jax import jax.numpy as jnp import matplotlib.pyplot as plt import numpy as np from euler_2d_double_shear_layer import ( # noqa: E501 DOM_T, DT, LOSSES_HISTORY_FILE, N_GRID_PLOT, NT, STATE_CKPT_SCRIPTNAME, X_MAX, X_MIN, Y_MAX, Y_MIN, build_space_state, state_checkpoint_postfix, ) from mpl_toolkits.axes_grid1 import make_axes_locatable from tqdm import tqdm FIG_DIR = os.path.join( os.path.dirname(os.path.abspath(__file__)), "fig", "double_shear_layer" ) os.makedirs(FIG_DIR, exist_ok=True) # Snapshots (PNG frames) are only rendered up to this time, not all the way to # DOM_T[1] -- the video only needs the roll-up phase, not the (visually static) # tail end of the run. T_SNAPSHOT_MAX = 11.0 NT_SNAPSHOT = int(round((T_SNAPSHOT_MAX - DOM_T[0]) / DT)) # Fixed color range/ticks shared by every rendering of omega, for a consistent # color scale across the whole figure set. FIELD_CLIM = 4.0 FIELD_CBAR_TICKS = (-4, -2, 0, 2, 4) CMAP = "managua" def load_space_state(frame_idx: int): """Rebuild the (untrained) state space and load its checkpoint at frame_idx.""" space_state = build_space_state(jax.random.PRNGKey(0)) success, space_state = space_state.load( STATE_CKPT_SCRIPTNAME, postfix=state_checkpoint_postfix(frame_idx) ) if not success: raise FileNotFoundError( f"No checkpoint found for frame {frame_idx} (scriptname=" f"{STATE_CKPT_SCRIPTNAME!r}, postfix={state_checkpoint_postfix(frame_idx)!r}). " "Run euler_2d_double_shear_layer.py first." ) return space_state def _omega_on_grid( space_state, x_grid: jnp.ndarray, y_grid: jnp.ndarray ) -> jnp.ndarray: """Evaluate omega (component 0 of the packed state) on a tensor grid.""" w = space_state.create_variables()[0] omega_batched = w.component(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_state, xy) return omega_vals.reshape(x_grid.shape[0], y_grid.shape[0]) def compute_enstrophy(space_state, n_grid: int = 128) -> jnp.ndarray: """Compute 0.5 * int_Omega omega^2 dx dy on a regular grid. The standard 2D-Euler diagnostic (conserved by the exact solution): a growing discretization/training error shows up as visible drift here. """ grid = jnp.linspace(X_MIN, X_MAX, n_grid, endpoint=False) omega_vals = _omega_on_grid(space_state, grid, grid) dx = (X_MAX - X_MIN) / n_grid dy = (Y_MAX - Y_MIN) / n_grid return 0.5 * jnp.sum(omega_vals**2) * dx * dy def compute_constraint_violation(space_state, n_grid: int = 128) -> jnp.ndarray: """Compute max |Delta psi + omega| on a regular grid -- see module docstring.""" w = space_state.create_variables()[0] omega, psi = w.component(0), w.component(1) residual = psi.laplacian("x") + omega residual_batched = residual.vmap_on_physical_variables() grid = jnp.linspace(X_MIN, X_MAX, n_grid, endpoint=False) xx, yy = jnp.meshgrid(grid, grid, indexing="ij") xy = jnp.stack([xx.ravel(), yy.ravel()], axis=-1) return jnp.max(jnp.abs(residual_batched(space_state, xy))) def save_omega_snapshot(space_state, frame_idx: int, t: float) -> None: """Save a single PNG snapshot of omega, for later assembly into a video.""" grid = jnp.linspace(X_MIN, X_MAX, N_GRID_PLOT, endpoint=False) omega_vals = jax.device_get(_omega_on_grid(space_state, grid, grid)) omega_vals = np.clip(omega_vals, -FIELD_CLIM, FIELD_CLIM) xx, yy = np.meshgrid(grid, grid) fig, ax = plt.subplots(figsize=(5, 5)) im = ax.contourf( xx, yy, omega_vals.T, levels=np.linspace(-FIELD_CLIM, FIELD_CLIM, 513), cmap=CMAP, extend="neither", ) im.set_edgecolor("face") ax.set_title(f"double shear layer (discrete PINN): vorticity, t={t:.3f}") ax.set_xlabel("x") ax.set_ylabel("y") ax.set_aspect("equal") divider = make_axes_locatable(ax) cax = divider.append_axes("right", size="5%", pad=0.05) fig.colorbar( im, cax=cax, extend="neither", label=r"$\omega$", ticks=FIELD_CBAR_TICKS ) fig.tight_layout() fig.savefig(os.path.join(FIG_DIR, f"omega_{frame_idx:05d}.png"), dpi=150) plt.close(fig) def plot_loss_history(losses_history: np.ndarray) -> None: """Plot the training loss history: initial condition fit + every macro step.""" fig, ax = plt.subplots(1, 1, figsize=(6, 4)) ax.semilogy(losses_history, color="tomato", lw=0.8) ax.set_xlabel("training epoch (initial fit, then every macro step)") ax.set_ylabel("total loss") ax.set_title("double shear layer (discrete PINN): training loss history") plt.tight_layout() plt.savefig(os.path.join(FIG_DIR, "double_shear_layer_loss_history.pdf")) plt.show() def plot_enstrophy_history(times: jnp.ndarray, enstrophy: jnp.ndarray) -> None: """Plot the enstrophy 0.5 * int omega^2 dx dy as a function of time.""" fig, ax = plt.subplots(1, 1, figsize=(6, 4)) ax.plot(times, enstrophy, color="tomato") ax.set_xlabel("t") ax.set_ylabel(r"$\frac{1}{2}\int \omega^2 \, dx \, dy$") ax.set_title("double shear layer (discrete PINN): enstrophy") plt.tight_layout() plt.savefig(os.path.join(FIG_DIR, "double_shear_layer_enstrophy.pdf")) plt.show() def plot_constraint_violation_history( times: jnp.ndarray, violation: jnp.ndarray ) -> None: """Plot max |Delta psi + omega| as a function of time -- see module docstring.""" fig, ax = plt.subplots(1, 1, figsize=(6, 4)) ax.semilogy(times, violation, color="tomato") ax.set_xlabel("t") ax.set_ylabel(r"$\max |\Delta \psi + \omega|$") ax.set_title("double shear layer (discrete PINN): DAE constraint violation") plt.tight_layout() plt.savefig(os.path.join(FIG_DIR, "double_shear_layer_constraint_violation.pdf")) plt.show() # %% if __name__ == "__main__": """Reload every state checkpoint and make the double shear layer diagnostics.""" print( f"Loading {NT + 1} state checkpoints (scriptname={STATE_CKPT_SCRIPTNAME!r})..." ) for old_png in glob.glob(os.path.join(FIG_DIR, "omega_*.png")): os.remove(old_png) enstrophy_history = [] constraint_violation_history = [] for frame_idx in tqdm(range(NT + 1)): space_state = load_space_state(frame_idx) enstrophy_history.append(float(compute_enstrophy(space_state))) constraint_violation_history.append( float(compute_constraint_violation(space_state)) ) if frame_idx <= NT_SNAPSHOT: save_omega_snapshot(space_state, frame_idx, DOM_T[0] + frame_idx * DT) times_history = DOM_T[0] + jnp.arange(NT + 1) * DT print("\nPlotting diagnostics...") plot_enstrophy_history(times_history, jnp.array(enstrophy_history)) plot_constraint_violation_history( times_history, jnp.array(constraint_violation_history) ) if os.path.exists(LOSSES_HISTORY_FILE): plot_loss_history(np.load(LOSSES_HISTORY_FILE)) else: print(f"No loss history found at {LOSSES_HISTORY_FILE}; skipping that plot.") print(f"\n{min(NT_SNAPSHOT, NT) + 1} omega snapshots saved to {FIG_DIR}") print("Assemble a video with e.g.:") print( f" ffmpeg -framerate 30 -i {FIG_DIR}/omega_%05d.png " "-pix_fmt yuv420p double_shear_layer_discrete_pinn.mp4" ) plt.show() # %%