r"""Verifies "heun"'s higher weak order against "directionwise"/"euler_maruyama". `NeuralSemiLagrangian`'s `kind_diffusion == "heun"` is documented (see `Characteristic._compute_feet_heun`) to reach weak order 2 in `dt` for additive-noise advection-diffusion SDEs, versus "euler_maruyama"'s weak order 1 -- *provided the advection field is not itself constant* (for constant advection, "heun"'s predictor and corrector coincide and it collapses to sub-stepped "euler_maruyama"; `advection_diffusion_2d.py`'s comparison uses constant advection, so it cannot show this gap at all). This example is built specifically to show it. .. math:: \partial_t u + a(x) \, \partial_x u - \sigma \, \partial_{xx} u = 0, \qquad a(x) = -K x a linear, mean-*expanding* drift (an Ornstein-Uhlenbeck-type advection field, but note the sign: see "A subtlety" below). Because `a` is linear and the diffusion is additive, this PDE is exactly solvable in closed form for a Gaussian initial condition -- essential here, since measuring a numerical scheme's convergence *order* requires an exact reference, not just "the error got smaller". A subtlety worth stating plainly, since it is easy to get backwards (an earlier version of this example did): this equation is the *backward Kolmogorov* equation for `u`, not the Fokker-Planck equation for a probability density -- `u(t, x) = E[u_0(Y_t) \mid Y_0 = x]` for the SDE `dY_s = -a(Y_s) ds + \sqrt{2 \sigma} dW_s` (note the *minus* sign on `a`, inherited from this module's `du/dt + a . grad u - sigma Delta u = 0` convention). With `a(x) = -K x`, that SDE has drift `+K Y_s ds`: it is *expanding*, not mean-reverting, so its conditional variance `(sigma/K)(exp(2 K t) - 1)` *grows* with `t`. Consequently `u`'s amplitude decays as `exp(-K t)` on top of the usual heat-kernel `sqrt(var_0 / var(t))` factor -- a Fokker-Planck-style closed form (as used correctly for *constant* advection in `advection_diffusion_nd.py`/`advection_diffusion_1d_1v.py`, where the two equations coincide since `a` has no spatial derivative) would be missing that factor and silently give the wrong answer here. See `exact_sol` below for the full formula, and `mean_t`/`var_t`/`amp_t` for its building blocks -- checked against an independent finite-difference solve of the PDE during development. This example has two parts: 1. A **direct, untrained** verification of the local (one-step) weak convergence order: `Characteristic.compute_feet` is called directly (no neural network, no training loop) to evaluate `E[u_0(foot)]` for one step of size `dt`, compared against the exact one-step solution. This is the cleanest possible test -- it isolates the SDE-integration scheme itself from every other source of error (network capacity, optimizer noise, collocation sampling) that would otherwise confound the measurement (as it did in an earlier attempt at this example: with training in the loop, "heun" and "euler_maruyama"'s finite Monte-Carlo sample budget dominates long before their underlying integration order becomes visible). Since there is no training here, the Monte-Carlo sample count (`N_DIFFUSION_SAMPLES`) can be made very large very cheaply. 2. A **light, trained** `NeuralSemiLagrangian.solve` run at a couple of `nt` values, to confirm the same ordering (`heun` more accurate than `euler_maruyama`, itself usually comparable to or better than `directionwise` despite the latter using the *exact* characteristic foot) survives in the full, trained pipeline -- more realistic, but noisier, for the reasons 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_1d import Segment1D 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 ( Characteristic, NeuralSemiLagrangian, ) K = 1.5 # advection field a(x) = -K*x; see the module docstring for its sign SIGMA = 0.1 # diffusion coefficient MU0 = 1.0 # initial bump center VAR0 = 0.08 # initial bump variance KINDS = ["directionwise", "euler_maruyama", "heun"] # --- Part 1: direct, untrained weak-order verification --------------------- DTS = [0.2, 0.1, 0.05, 0.025, 0.0125] N_DIFFUSION_SAMPLES = 500_000 # cheap here: no training loop, just one MC average X0_PROBE = jnp.linspace(-0.5, 2.0, 20).reshape(-1, 1) # several starting points # --- Part 2: light, trained NSL comparison ---------------------------------- X_BOUNDS = (-10.0, 10.0) # generous margin: the backward foot expands (see # the module docstring), so a tight, non-periodic domain would let bulk # collocation points map outside the trained region every step -- the same # "wide margin" pattern used in advection_diffusion_1d_1v.py, and, unlike # there, load-bearing here (an earlier version of this example, without the # margin, gave a nonsensical, non-decreasing error irrespective of the # exact-solution bug described above). DOMAIN_BOUNDS = ((X_BOUNDS[0],), (X_BOUNDS[1],)) DOM_X = Segment1D(X_BOUNDS, is_main_domain=True) SAMPLER = TensorizedSampler([DomainSampler(DOM_X)], model_type="x") END_TIME = 0.6 DOM_T = (0.0, END_TIME) N_COLLOC = 4000 N_EPOCHS_INIT = 300 N_EPOCHS = 30 NT_LIST_TRAINED = [4, 8, 16] N_DIFFUSION_SAMPLES_TRAINED = 40 N_RK_STEPS_TRAINED = 4 SEED = 0 def mean_t(t: jnp.ndarray) -> jnp.ndarray: """Mean of the exact solution's Gaussian, `mu0 * exp(-K t)`.""" return MU0 * jnp.exp(-K * t) def var_t(t: jnp.ndarray) -> jnp.ndarray: """Variance of the exact solution's Gaussian. `var0 * exp(-2 K t) + (sigma/K) * (1 - exp(-2 K t))`: the same formula a (wrong, for this equation) Fokker-Planck derivation would also produce -- only `amp_t` differs. Checked against a finite-difference solve. """ return VAR0 * jnp.exp(-2 * K * t) + (SIGMA / K) * (1 - jnp.exp(-2 * K * t)) def amp_t(t: jnp.ndarray) -> jnp.ndarray: """Amplitude of the exact solution: heat-kernel factor times `exp(-K t)`. The `exp(-K t)` is the piece a Fokker-Planck-style guess would miss; see the module docstring's "A subtlety" section. """ return jnp.sqrt(VAR0 / var_t(t)) * jnp.exp(-K * t) def u0(y: jnp.ndarray) -> jnp.ndarray: """Initial condition (also the test function probed in Part 1).""" return jnp.exp(-((y - MU0) ** 2) / (2 * VAR0)) def exact_sol(t: jnp.ndarray, x: jnp.ndarray) -> jnp.ndarray: """Exact solution `u(t, x) = amp(t) * exp(-(x - mean(t))^2 / (2 var(t)))`.""" x = jnp.atleast_2d(x) t = jnp.atleast_1d(t).reshape(x.shape[0], 1)[:, 0] result = amp_t(t) * jnp.exp(-((x[:, 0] - mean_t(t)) ** 2) / (2 * var_t(t))) 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) = -K * x -- non-constant, unlike every other example in this directory (required: constant advection makes "heun" collapse to "euler_maruyama"; see the module docstring).""" return -K * x def exact_characteristic_foot(t: jnp.ndarray, x: jnp.ndarray, dt: float) -> jnp.ndarray: """Exact backward characteristic foot for `dx/ds = -K x`: `x * exp(K dt)`.""" return x * jnp.exp(K * dt) def local_weak_error(kind: str, dt: float, n_diffusion_samples: int | None) -> float: """One-step weak error `|E[u0(foot)] - u_exact(dt, x0)|`, RMS over `X0_PROBE`. Calls `Characteristic.compute_feet` directly -- no network, no training -- so `n_diffusion_samples` can be pushed far higher than would ever be affordable inside a training loop. """ characteristic = Characteristic( advection_field=advection_field, dt=dt, model_type="x", exact_characteristic_foot=exact_characteristic_foot if kind != "heun" else None, diffusion_coefficient_x=SIGMA, kind_diffusion=kind, n_diffusion_samples=n_diffusion_samples, ) t0 = jnp.zeros_like(X0_PROBE) feet = characteristic.compute_feet(t0, X0_PROBE) u_avg = sum(u0(foot[0]) for foot in feet) / len(feet) t_dt = jnp.full_like(X0_PROBE, dt) return float(jnp.sqrt(jnp.mean((u_avg - exact_sol(t_dt, X0_PROBE)) ** 2))) def fitted_order(dts: list[float], errors: list[float]) -> float: """Least-squares slope of log(error) vs log(dt): the observed local order.""" slope, _ = np.polyfit(np.log(dts), np.log(errors), 1) return float(slope) def make_trained_solver(kind: str, nt: int) -> NeuralSemiLagrangian: """Build an (untrained) NSL solver for Part 2, given `kind` and `nt`. `N_DIFFUSION_SAMPLES_TRAINED` (well above `"directionwise"`'s default `2 * dim = 2` for this 1D problem) and `N_RK_STEPS_TRAINED` keep "euler_maruyama"/"heun"'s Monte-Carlo noise from swamping this comparison, the same issue Part 1's module docstring describes -- without them, both stochastic kinds' errors here are dominated by noise rather than by the ordering this part is meant to illustrate. """ return NeuralSemiLagrangian( main_domain=DOM_X, time_domain=DOM_T, sampler=SAMPLER, dt=(DOM_T[1] - DOM_T[0]) / nt, 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=SIGMA, kind_diffusion=kind, n_diffusion_samples=N_DIFFUSION_SAMPLES_TRAINED if kind != "directionwise" else None, n_rk_steps=N_RK_STEPS_TRAINED if kind == "heun" else 1, ) if __name__ == "__main__": # ------------------------------------------------------------------ # Part 1: direct, untrained weak-order verification. # ------------------------------------------------------------------ print("Part 1: local (one-step) weak-error verification (no training)...") local_errors: dict[str, list[float]] = {} for kind in KINDS: start = time.perf_counter() n_samples = N_DIFFUSION_SAMPLES if kind != "directionwise" else None errors = [local_weak_error(kind, dt, n_samples) for dt in DTS] local_errors[kind] = errors order = fitted_order(DTS, errors) errs_str = ", ".join(f"{e:.2e}" for e in errors) print( f" {kind:>15s}: errors = [{errs_str}] " f"(fitted local order: {order:.2f}, {time.perf_counter() - start:.1f}s)" ) print( " Theory (additive noise, Kloeden & Platen): local order 2 for " "'directionwise'/'euler_maruyama' (weak order 1), local order 3 for " "'heun' (weak order 2). 'heun' should be markedly more accurate than " "both at every dt above, with a fitted order exceeding theirs -- how " "close to the asymptotic '3' depends on how far the finite " "N_DIFFUSION_SAMPLES Monte-Carlo floor lets the smallest dt values " "resolve it (see the module docstring).\n" ) # ------------------------------------------------------------------ # Part 2: light, trained NSL comparison. # ------------------------------------------------------------------ print("Part 2: light, trained NeuralSemiLagrangian comparison...") key = jax.random.PRNGKey(SEED) nn = MLP(in_size=1, out_size=1, hidden_sizes=[16, 16, 16], key=key) space = ApproximationSpace({"x": 1}, [(nn, "scalar", None)], model_type="x") key, init_solver = make_trained_solver(KINDS[0], nt=NT_LIST_TRAINED[0]).initialize( key, space, f_init, N_EPOCHS_INIT, N_COLLOC ) init_space = init_solver.space trained_errors: dict[str, list[float]] = {kind: [] for kind in KINDS} finest_solve: dict[str, NeuralSemiLagrangian] = {} for kind in KINDS: for nt in NT_LIST_TRAINED: key, neural_sl = make_trained_solver(kind, nt).solve( key, init_space, N_EPOCHS, N_COLLOC ) trained_errors[kind].append(float(neural_sl.errors_over_time[-1, 0])) if nt == NT_LIST_TRAINED[-1]: finest_solve[kind] = neural_sl errs_str = ", ".join(f"{e:.2e}" for e in trained_errors[kind]) print(f" {kind:>15s} (nt={NT_LIST_TRAINED}): errors = [{errs_str}]") # ------------------------------------------------------------------ # Plot 1: Part 1's log-log local weak error, with order 2/3 references. # ------------------------------------------------------------------ fig1, ax1 = plt.subplots(figsize=(6, 5)) for kind in KINDS: ax1.loglog(DTS, local_errors[kind], marker="o", label=kind) dt_ref = np.array([DTS[0], DTS[-1]]) for order, style in [(2, "k--"), (3, "k:")]: err_ref = local_errors["directionwise"][0] * (dt_ref / DTS[0]) ** order ax1.loglog(dt_ref, err_ref, style, label=f"order {order} (reference)") ax1.set_xlabel("dt") ax1.set_ylabel("local (one-step) weak error") ax1.set_title("Local weak-convergence order (untrained)") ax1.legend() fig1.tight_layout() # ------------------------------------------------------------------ # Plot 2: Part 2's trained-solver error vs nt. # ------------------------------------------------------------------ fig2, ax2 = plt.subplots(figsize=(6, 5)) for kind in KINDS: dts_trained = [END_TIME / nt for nt in NT_LIST_TRAINED] ax2.loglog(dts_trained, trained_errors[kind], marker="o", label=kind) ax2.set_xlabel("dt") ax2.set_ylabel("final relative L2 error") ax2.set_title("Trained NeuralSemiLagrangian: same ordering, noisier") ax2.legend() fig2.tight_layout() # ------------------------------------------------------------------ # Plot 3: final 1D profile (Part 2, nt=NT_LIST_TRAINED[-1]) vs exact. # ------------------------------------------------------------------ xg = jnp.linspace(*X_BOUNDS, 500).reshape(-1, 1) t_final = jnp.full((xg.shape[0], 1), END_TIME) u_exact_final = jax.device_get(exact_sol(t_final, xg)).ravel() fig3, ax3 = plt.subplots(figsize=(7, 5)) ax3.plot(xg.ravel(), u_exact_final, "k-", label="exact", linewidth=2) for kind in KINDS: neural_sl = finest_solve[kind] variables = neural_sl.space.create_variables() all_together = ParamVecFunction.cat(variables) batched_func = all_together.vmap_on_physical_variables() u_pred = jax.device_get(batched_func(neural_sl.space, xg)).ravel() ax3.plot(xg.ravel(), u_pred, "--", label=kind) ax3.set_xlim(-1.0, 2.5) ax3.set_xlabel("x") ax3.set_ylabel("u") ax3.set_title(f"Final profile at t={END_TIME} (nt={NT_LIST_TRAINED[-1]})") ax3.legend() fig3.tight_layout() plt.show()