r"""Compares NSL's diffusion strategies on a 2D advection-diffusion equation. .. math:: \partial_t u + a \cdot \nabla u - \sigma \Delta u & = 0 \quad \text{in } \Omega \times (0, T) \\ u & = u_0 \quad \text{on } \Omega \times \{0\} \\ u & \text{ periodic on } \partial \Omega where :math:`u: \Omega \times (0, T) \to \mathbb{R}` is the unknown function, :math:`\Omega = (-1, 1)^2 \subset \mathbb{R}^2` is the spatial domain, :math:`(0, T) \subset \mathbb{R}` is the time domain, :math:`a` is a constant advection velocity and :math:`\sigma` is a constant, isotropic diffusion coefficient. The physical constants and exact solution are the DIM=2 case of `advection_diffusion_nd.py` (see that module for the derivation), kept unchanged here so results are directly comparable. `Characteristic.compute_feet` advances diffusion by averaging `u^n` over a small stencil around the (purely advective) characteristic foot at every time step; `kind_diffusion` selects how that stencil is built: - "directionwise", "simplex", "hypercube": deterministic stencils (fixed unit directions, scaled and averaged so their second moment reproduces the Laplacian to leading order; see `_make_diffusion_directions`). - "euler_maruyama": genuine Gaussian increments -- the Euler-Maruyama discretization of the diffusion SDE `dX = sqrt(2 sigma) dW` underlying the heat equation (see `_euler_maruyama_offsets`), added *after* a purely advective (RK4 or exact) foot. Unlike the deterministic stencils (one fixed translation shared by the whole batch, per direction), every collocation point draws its own independent increment, and the average is a Monte-Carlo estimate rather than an exact finite-difference identity. - "heun": a higher-order stochastic solver -- rather than a post-hoc offset, it resolves advection and diffusion *together*, sub-step by sub-step, with a stochastic Heun (predictor-corrector) scheme (see `_compute_feet_heun`). Since diffusion here is additive noise (a constant coefficient), this reaches weak order 2 in `dt`, versus "euler_maruyama"'s weak order 1 -- a standard result for SDEs with additive noise (Kloeden & Platen) -- *when the advection field is not itself constant*. On this example's constant `VELOCITY`, "heun"'s predictor and corrector coincide and it collapses, step for step, to sub-stepped "euler_maruyama": comparing them here mainly demonstrates that collapse, not "heun"'s accuracy edge, which only shows up for spatially/temporally varying advection (see `test_heun_beats_plain_euler_for_nonconstant_advection` in the test suite for a case where it does). This example runs two comparisons: 1. All five strategies at a fixed number of time steps, with `euler_maruyama` and `heun` using their default sample count (`2 * dim`, matching "directionwise"'s cost) for a fair per-step comparison: final errors, wall-clock time, error over time, and the final 2D solution field. 2. A time-convergence study (halving `dt` a few times) that fits an observed convergence order to each strategy's error, to confirm the scheme is (weak) first-order in time -- as it must be, since every diffusion strategy here is one explicit-Euler step per time step, whether the diffusion part of that step is a deterministic finite-difference-like average or a Monte-Carlo one. "directionwise" and "hypercube" reproduce the Laplacian *exactly* (up to network-fitting error) at every step -- both stencils are invariant under `v -> -v` (every direction's antipode is in the set too), which forces their third moment to vanish exactly and leaves no `O(dt^1.5)` term in the per-step truncation error, only `O(dt^2)`; their truncation error dominates at any practical `dt`, and the observed order should land close to 1 (see `tests/.../ test_advection_diffusion_2d_error_decreases_with_more_time_steps` for a regression-tested version of this same check on "directionwise"). "simplex" is *not* antipodally symmetric, so it keeps a nonzero third moment and an uncancelled `O(dt^1.5)` local term: its observed order is close to 0.5, not 1, and no amount of extra training or sampling closes that gap -- see `_make_diffusion_directions`'s docstring for the full derivation. "euler_maruyama" is different again: each step's diffusion estimate also carries Monte-Carlo *variance*, from a finite sample budget (`N_DIFFUSION_SAMPLES_CONV` below, deliberately larger here than the default so the O(dt) trend has a chance to show through at all) that does not automatically shrink as `dt` shrinks. If that variance floor is comparable to (or above) the truncation error at the finest `dt` tested, the observed order will fall short of 1 (it can even go negative -- error *growing* with more, smaller time steps, as the accumulated per-step noise outpaces the shrinking per-step bias) -- not a bug, but the expected signature of an under-sampled Monte-Carlo estimator, and the reason "euler_maruyama" defaults to matching the deterministic stencils' *cost* rather than their accuracy. Reaching the same order at the same `dt` would require growing the sample budget (`n_diffusion_samples`) or the epoch count as `dt` shrinks, which is left to the user rather than hard-coded here. "heun" carries the same finite-sample-budget caveat -- on this example's *constant* advection field it collapses to "euler_maruyama", so it inherits the same Monte-Carlo-dominated, order-short-of-1 behavior, not an improved one: its weak-order-2 advantage is about resolving a *non-constant* drift more accurately, which this test problem doesn't exercise (see the module docstring above). """ import time import jax import jax.numpy as jnp import matplotlib.pyplot as plt import numpy as np from scimba_jax.domains.meshless_domains.domains_nd import HypercubeND from scimba_jax.nonlinear_approximation.approximation_spaces.approximation_spaces import ( # noqa: E501 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.neural_semi_lagrangian import ( NeuralSemiLagrangian, ) DIM = 2 N_COLLOC = 2000 * DIM N_EPOCHS_INIT = 300 N_EPOCHS = 30 NT = 16 # number of time steps used for the main (fixed-nt) comparison SEED = 0 DIFFUSION_COEFFICIENT = 0.1 VELOCITY = 1.0 # advection speed, identical in every direction FREQUENCY = jnp.pi DOMAIN_BOUNDS = ((-1.0,), (1.0,)) # isotropic: same bounds along every axis DOM_X = HypercubeND([(-1.0, 1.0)] * DIM, is_main_domain=True) SAMPLER = TensorizedSampler([DomainSampler(DOM_X)], model_type="x") # End time set so that the diffusive decay halves the amplitude of the # solution's oscillating part, as in advection_diffusion_nd.py. END_TIME = jnp.log(2.0) / (DIM * DIFFUSION_COEFFICIENT * FREQUENCY**2) DOM_T = (0.0, float(END_TIME)) DT = (DOM_T[1] - DOM_T[0]) / NT # Phase shift per axis, so that no two axes carry the same phase. PHASE_SHIFT = jnp.linspace(0, 1 - 1 / DIM, DIM) # The diffusion strategies compared, in plotting order. KINDS = ["directionwise", "simplex", "hypercube", "euler_maruyama", "heun"] # Time-convergence study: number of time steps tried, and the (larger, see # module docstring) Monte-Carlo sample count used only for that study, for # both "euler_maruyama" and "heun". NT_LIST = [8, 16, 32] N_DIFFUSION_SAMPLES_CONV = 40 def exact_sol(t: jnp.ndarray, x: jnp.ndarray) -> jnp.ndarray: r"""Exact solution: a single advecting/diffusing Fourier mode. See `advection_diffusion_nd.py`'s `exact_sol` for the derivation; ported unchanged (same constants) so the diffusion strategies below are compared on identical footing. """ x = jnp.atleast_2d(x) t = jnp.atleast_1d(t).reshape(x.shape[0], 1)[:, 0] arg = jnp.sum(x - PHASE_SHIFT[None, :] - VELOCITY * t[:, None], axis=1) decay = jnp.exp(-DIM * DIFFUSION_COEFFICIENT * FREQUENCY**2 * t) result = 2 + jnp.sin(FREQUENCY * arg) * decay return result if result.shape[0] == 1 else result.reshape(-1, 1) def f_init(x: jnp.ndarray) -> jnp.ndarray: """Initial condition function.""" t = jnp.zeros((1,)) if x.ndim == 1 else jnp.zeros((x.shape[0], 1)) return exact_sol(t, x) def advection_field(t: jnp.ndarray, x: jnp.ndarray) -> jnp.ndarray: """Advection field a(t, x) = VELOCITY, identical in every direction.""" return jnp.ones_like(x) * VELOCITY def exact_characteristic_foot(t: jnp.ndarray, x: jnp.ndarray, dt: float) -> jnp.ndarray: """Exact characteristic foot for constant advection: X(t) = x - a * dt.""" return x - VELOCITY * dt def make_solver( kind_diffusion: str, dt: float = DT, n_diffusion_samples: int | None = None, ) -> NeuralSemiLagrangian: """Build an (untrained) NSL solver for the given diffusion strategy.""" return NeuralSemiLagrangian( main_domain=DOM_X, time_domain=DOM_T, sampler=SAMPLER, dt=dt, advection_field=advection_field, exact_characteristic_foot=exact_characteristic_foot, periodic=True, domain_bounds=DOMAIN_BOUNDS, out_size=1, model_type="x", exact_solution=exact_sol, diffusion_coefficient_x=DIFFUSION_COEFFICIENT, kind_diffusion=kind_diffusion, n_diffusion_samples=n_diffusion_samples, ) def observed_order(dts: list[float], errors: list[float]) -> float: """Least-squares slope of log(error) vs log(dt): the observed convergence order.""" slope, _ = np.polyfit(np.log(dts), np.log(errors), 1) return float(slope) if __name__ == "__main__": key = jax.random.PRNGKey(SEED) out_size = 1 hidden_sizes = [7 * DIM] * 3 nn = MLP(in_size=DIM, out_size=out_size, hidden_sizes=hidden_sizes, key=key) space = ApproximationSpace({"x": DIM}, [(nn, "scalar", None)], model_type="x") # Fit the initial condition once (diffusion-strategy- and dt-independent: # `initialize` never touches `dt`) and reuse it everywhere below, so every # comparison isolates the diffusion approximation / time discretization # from initialization noise. print("Fitting the shared initial condition...") start_init = time.perf_counter() key, init_solver = make_solver(KINDS[0]).initialize( key, space, f_init, N_EPOCHS_INIT, N_COLLOC, file_name="nsl_ad_2d_init", retrain=False, ) init_space = init_solver.space print(f"... done in {time.perf_counter() - start_init:.2f} seconds\n") # ------------------------------------------------------------------ # Part 1: all four strategies at a fixed NT, matched per-step cost. # ------------------------------------------------------------------ results = {} for kind in KINDS: print(f"Solving with kind_diffusion={kind!r} (nt={NT})...") start_solve = time.perf_counter() key, neural_sl = make_solver(kind).solve(key, init_space, N_EPOCHS, N_COLLOC) elapsed = time.perf_counter() - start_solve errors = neural_sl.errors_over_time print( f" ... done in {elapsed:.2f}s, final relative L2 error: " f"{errors[-1, 0]:.2e}, final relative Linf error: {errors[-1, 1]:.2e}\n" ) results[kind] = (neural_sl, elapsed) # ------------------------------------------------------------------ # Part 2: time-convergence study (see the module docstring for why # "euler_maruyama" and "heun" need a larger, dedicated sample count here). # ------------------------------------------------------------------ print("Time-convergence study (halving dt, tracking the observed order)...") # Expected order per strategy -- see the module docstring and # `_make_diffusion_directions` for the derivation: "directionwise" and # "hypercube" are antipodally symmetric (order 1); "simplex" is not, and # its uncancelled third moment caps it at order 0.5; "euler_maruyama"'s # and "heun"'s order-1/2 weak convergence only shows once their # Monte-Carlo variance is driven below the truncation error, which the # modest `N_DIFFUSION_SAMPLES_CONV` used here does not guarantee -- and, # for "heun" specifically, only once the advection field is non-constant # (not the case in this example; see the module docstring). mc_caveat = "~1-2 once oversampled enough; not guaranteed here" expected_order = { "directionwise": 1.0, "hypercube": 1.0, "simplex": 0.5, "euler_maruyama": mc_caveat, "heun": mc_caveat, } conv_errors: dict[str, list[float]] = {kind: [] for kind in KINDS} conv_dts = [(DOM_T[1] - DOM_T[0]) / nt for nt in NT_LIST] for kind in KINDS: n_samples = ( N_DIFFUSION_SAMPLES_CONV if kind in ("euler_maruyama", "heun") else None ) for nt, dt in zip(NT_LIST, conv_dts): key, neural_sl = make_solver(kind, dt, n_samples).solve( key, init_space, N_EPOCHS, N_COLLOC ) conv_errors[kind].append(float(neural_sl.errors_over_time[-1, 0])) order = observed_order(conv_dts, conv_errors[kind]) errs_str = ", ".join(f"{e:.2e}" for e in conv_errors[kind]) print( f" {kind:>15s}: errors = [{errs_str}] (observed order: {order:.2f}, " f"expected: {expected_order[kind]})" ) print() # ------------------------------------------------------------------ # Plot 1: relative L2 error over time (Part 1, fixed NT), all strategies. # ------------------------------------------------------------------ times_plot = jnp.arange(NT + 1) * DT + DOM_T[0] fig_err, ax_err = plt.subplots(figsize=(6, 5)) for kind in KINDS: neural_sl, _ = results[kind] ax_err.semilogy( jax.device_get(times_plot), jax.device_get(neural_sl.errors_over_time[:, 0]), marker="o", label=kind, ) ax_err.set_xlabel("time") ax_err.set_ylabel("relative L2 error") ax_err.set_title(f"NSL diffusion strategies: error over time (nt={NT})") ax_err.legend() fig_err.tight_layout() # ------------------------------------------------------------------ # Plot 2: time-convergence study, log-log, with an order-1 reference. # ------------------------------------------------------------------ fig_conv, ax_conv = plt.subplots(figsize=(6, 5)) for kind in KINDS: ax_conv.loglog(conv_dts, conv_errors[kind], marker="o", label=kind) dt_ref = np.array([conv_dts[0], conv_dts[-1]]) err_ref = conv_errors["hypercube"][0] * (dt_ref / conv_dts[0]) ax_conv.loglog(dt_ref, err_ref, "k--", label="order 1 (reference)") ax_conv.set_xlabel("dt") ax_conv.set_ylabel("final relative L2 error") ax_conv.set_title("NSL diffusion strategies: time-convergence order") ax_conv.legend() fig_conv.tight_layout() # ------------------------------------------------------------------ # Plot 3: final 2D solution field for every strategy (Part 1, nt=NT). # ------------------------------------------------------------------ n_grid = 200 grid_1d = jnp.linspace(-1.0, 1.0, n_grid) x1, x2 = jnp.meshgrid(grid_1d, grid_1d) x_plot = jnp.stack([x1.ravel(), x2.ravel()], axis=1) t_final = jnp.full((x_plot.shape[0], 1), DOM_T[-1]) u_exact = jax.device_get(exact_sol(t_final, x_plot)).reshape(n_grid, n_grid) fig, axes = plt.subplots(1, len(KINDS) + 1, figsize=(4 * (len(KINDS) + 1), 4)) cf = axes[0].contourf(x1, x2, u_exact, cmap="turbo", levels=100) plt.colorbar(cf, ax=axes[0]) axes[0].set_title(f"Exact solution at t={DOM_T[-1]:.3f}") axes[0].set_aspect("equal") for ax, kind in zip(axes[1:], KINDS): neural_sl, elapsed = results[kind] space = neural_sl.space variables = space.create_variables() all_together = ParamVecFunction.cat(variables) batched_func = all_together.vmap_on_physical_variables() u_pred = jax.device_get(batched_func(space, x_plot)).reshape(n_grid, n_grid) cf = ax.contourf(x1, x2, u_pred, cmap="turbo", levels=100) plt.colorbar(cf, ax=ax) final_l2 = neural_sl.errors_over_time[-1, 0] ax.set_title(f"{kind}\n(rel. L2: {final_l2:.2e}, {elapsed:.1f}s)") ax.set_aspect("equal") fig.suptitle(f"{DIM}D advection-diffusion: comparing NSL diffusion strategies") fig.tight_layout() plt.show()