r"""Post-processing companion to ``euler_2d_double_shear_layer.py``. That script trains the double shear layer NSL solve and, at every macro time step, checkpoints the trained ``omega`` :class:`ApproximationSpace` to disk (via the generic ``ScimbaPytree.save`` facility) instead of rendering a figure inline -- the matplotlib rendering used to happen once per macro step, inside the training loop. This script does no training: it rebuilds an (untrained) omega 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, then plots the enstrophy history, plots 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. """ # %% 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, OMEGA_CKPT_SCRIPTNAME, X_MAX, X_MIN, Y_MAX, Y_MIN, _omega_on_grid, build_space_omega, omega_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 (video snapshots and # comparison PDFs alike), for a consistent color scale across the whole figure set. FIELD_CLIM = 4.0 FIELD_CBAR_TICKS = (-4, -2, 0, 2, 4) # omega develops finer filamentation as t grows (Kelvin-Helmholtz roll-up of the two # shear layers); past this time N_GRID_PLOT under-resolves it and visibly smooths out # fine structure in both the snapshot PNGs and the video, so those frames (and any # comparison PDF at t >= T_FINE_GRID_START) are rendered on a finer grid instead. T_FINE_GRID_START = 9.0 N_GRID_PLOT_FINE = 1024 N_CONTOUR = 512 + 1 def _plot_grid_n(t: float) -> int: """Grid resolution for a snapshot at time t -- finer past T_FINE_GRID_START.""" return N_GRID_PLOT_FINE if t >= T_FINE_GRID_START else N_GRID_PLOT def load_space_omega(frame_idx: int): """Rebuild the (untrained) omega space and load its checkpoint at frame_idx.""" space_omega = build_space_omega(jax.random.PRNGKey(0)) success, space_omega = space_omega.load( OMEGA_CKPT_SCRIPTNAME, postfix=omega_checkpoint_postfix(frame_idx) ) if not success: raise FileNotFoundError( f"No checkpoint found for frame {frame_idx} (scriptname=" f"{OMEGA_CKPT_SCRIPTNAME!r}, postfix={omega_checkpoint_postfix(frame_idx)!r}). " "Run euler_2d_double_shear_layer.py first." ) return space_omega def compute_enstrophy(space_omega, 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 NSL 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_omega, 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 save_omega_snapshot(space_omega, 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, _plot_grid_n(t), endpoint=False) omega_vals = jax.device_get(_omega_on_grid(space_omega, 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, N_CONTOUR), cmap=COMPARISON_CMAPS[0], extend="neither", ) im.set_edgecolor("face") ax.set_title(f"double shear layer: 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) # Qualitative comparison figures: t=6 (colorbar on the left) and t=10 (colorbar on # the right), each rendered with every colormap in COMPARISON_CMAPS. COMPARISON_TIMES = (6.0, 10.0) COMPARISON_CMAPS = ["managua"] COMPARISON_COLORBAR_SIDE = {6.0: "left", 10.0: "right"} # contourf with N_CONTOUR levels emits one vector polygon per level band, which bloats the # PDF to tens of MB per figure. Rasterizing just the filled field (zorder below the # ax.set_rasterization_zorder threshold) while keeping the colorbar vector fixes the # file size without blurring the (already annotation-free) colorbar text. COMPARISON_PDF_RASTER_DPI = 300 def save_omega_pdf(space_omega, t: float, cmap: str) -> None: """Save one annotation-free PDF snapshot of omega, for a single colormap. Only the colorbar carries any annotation (ticks, label) -- the field itself has no axes, ticks, or title, per COMPARISON_TIMES/COMPARISON_CMAPS above. The color range is fixed to [-FIELD_CLIM, FIELD_CLIM], same as the video snapshots. """ colorbar_side = COMPARISON_COLORBAR_SIDE[t] grid = jnp.linspace(X_MIN, X_MAX, _plot_grid_n(t), endpoint=False) omega_vals = jax.device_get(_omega_on_grid(space_omega, grid, grid)) vmin, vmax = -FIELD_CLIM, FIELD_CLIM omega_vals = np.clip(omega_vals, vmin, vmax) xx, yy = np.meshgrid(grid, grid) fig, ax = plt.subplots(figsize=(5, 5)) im = ax.contourf( xx, yy, omega_vals.T, levels=np.linspace(vmin, vmax, N_CONTOUR), cmap=cmap, extend="neither", zorder=-10, ) # contourf renders each level band as a separate filled polygon; antialiasing # then leaves thin seams visible at the shared edges between bands, most # noticeable with high-contrast colormaps like "managua". Matching the edge # color to the face color hides the seams while keeping contourf's filled # look (as opposed to switching to pcolormesh). im.set_edgecolor("face") # Rasterize only the filled field (everything at zorder < 0): the colorbar and # its ticks/label stay vector, the N_CONTOUR-band field becomes one bitmap layer. ax.set_rasterization_zorder(0) ax.set_aspect("equal") ax.axis("off") divider = make_axes_locatable(ax) cax = divider.append_axes(colorbar_side, size="5%", pad=0.05) fig.colorbar(im, cax=cax, extend="neither", ticks=FIELD_CBAR_TICKS) if colorbar_side == "left": cax.yaxis.set_ticks_position("left") cax.yaxis.set_label_position("left") fig.savefig( os.path.join(FIG_DIR, f"omega_t{t:g}_{cmap}.pdf"), bbox_inches="tight", dpi=COMPARISON_PDF_RASTER_DPI, ) plt.close(fig) def plot_loss_history(losses_history: np.ndarray) -> None: """Plot the training loss history: initial condition fit + every transport substep. Args: losses_history: 1D array of "total" loss values, in chronological order. """ 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 transport substep)") ax.set_ylabel("total loss") ax.set_title("double shear layer: NSL 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: enstrophy") plt.tight_layout() plt.savefig(os.path.join(FIG_DIR, "double_shear_layer_enstrophy.pdf")) plt.show() # %% if __name__ == "__main__": """Reload every omega checkpoint and make the double shear layer diagnostics.""" print( f"Loading {NT + 1} omega checkpoints (scriptname={OMEGA_CKPT_SCRIPTNAME!r})..." ) for old_png in glob.glob(os.path.join(FIG_DIR, "omega_*.png")): os.remove(old_png) comparison_frame_idx = { t: int(round((t - DOM_T[0]) / DT)) for t in COMPARISON_TIMES } comparison_spaces = {} enstrophy_history = [] for frame_idx in tqdm(range(NT + 1)): space_omega = load_space_omega(frame_idx) enstrophy_history.append(float(compute_enstrophy(space_omega))) if frame_idx <= NT_SNAPSHOT: save_omega_snapshot(space_omega, frame_idx, DOM_T[0] + frame_idx * DT) for t, t_frame_idx in comparison_frame_idx.items(): if frame_idx == t_frame_idx: comparison_spaces[t] = space_omega times_history = DOM_T[0] + jnp.arange(NT + 1) * DT print("\nPlotting diagnostics...") plot_enstrophy_history(times_history, jnp.array(enstrophy_history)) print("\nSaving comparison PDFs (t, colormap)...") for t, space_omega_t in comparison_spaces.items(): for cmap in COMPARISON_CMAPS: save_omega_pdf(space_omega_t, t, cmap) 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{NT_SNAPSHOT + 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.mp4" ) plt.show() # %%