r"""Verifies "strang"'s higher splitting order against "lie" for NSL reaction terms. .. math:: \partial_t u + c \, \partial_x u = s(t, x, u), \qquad s(t, x, u) = r(x) \, u \, (1 - u), \qquad r(x) = r_0 \, (1 + 0.5 \sin(2 \pi x)) Constant advection plus a nonlinear, *spatially-varying* logistic reaction, added via `NeuralSemiLagrangian`'s `reaction_term`/`kind_reaction` (see `ReactionStep` and the "Reaction" section of `neural_semi_lagrangian.py`'s module docstring for the underlying Lie/Strang operator splitting). The rate `r(x)` is deliberately *not* constant: a spatially-homogeneous rate's reaction flow is translation-invariant, so it commutes *exactly* with the pure-translation transport flow `T_dt`, making the Lie/Strang splitting error identically zero regardless of `dt` (verified during development: both kinds then reproduced the closed-form homogeneous-rate solution to floating-point precision, for every `dt`) -- which would hide any order difference. With a genuinely `x`-dependent rate, there is no closed form for this PDE, so both parts below check accuracy by comparison against numerically-computed references rather than an exact solution. This example has two parts: 1. A **direct, untrained** verification of the *local self-convergence order* of the splitting scheme itself: `Characteristic.compute_feet` and `ReactionStep.evolve` are composed by hand (mirroring `NeuralSemiLagrangian .solve`'s `target_func` exactly), with no network and no training, so the measured order reflects only the splitting scheme. Checked by self-convergence against a very fine "strang" reference (composed the same way, just with a tiny `dt`). 2. A **light, trained** `NeuralSemiLagrangian.solve` run, to confirm the full pipeline (network fitting, at each time step, to the reaction-split target) correctly reproduces the same PDE. Since Part 1's self-convergence reference is built from the *same* splitting family being tested (so it cannot, by itself, catch a bug shared by both "lie" and "strang", e.g. a wrong sign in `reaction_term`), this part instead checks against an **independent** reference: a classical method-of-lines discretization (spectral derivative in `x`, exact up to machine precision for this smooth periodic profile; plain RK4 in time) of the PDE directly -- no operator splitting at all -- built by `rk4_reference_solution`. """ 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 ( Characteristic, NeuralSemiLagrangian, ReactionStep, ) KINDS = ["lie", "strang"] C = 0.5 # constant advection velocity R0 = 3.0 # reaction-rate magnitude DOMAIN_BOUNDS = ((0.0,), (1.0,)) TOTAL_TIME = 0.4 # --- Part 1: direct, untrained self-convergence order check ---------------- PART1_NTS = [4, 8, 16, 32] PART1_REFERENCE_NT = 128 # fine "strang" self-convergence reference PART1_N_REACTION_SUBSTEPS = 8 X_PROBE = jnp.linspace(0.0, 1.0, 41).reshape(-1, 1) # --- Part 2: light, trained NSL comparison ---------------------------------- DOM_X = HypercubeND([(0.0, 1.0)], is_main_domain=True) DOM_T = (0.0, TOTAL_TIME) SAMPLER = TensorizedSampler([DomainSampler(DOM_X)], model_type="x") N_COLLOC = 1000 N_EPOCHS_INIT = 500 N_EPOCHS = 30 NT_LIST_TRAINED = [1, 2, 4, 8, 16] SEED = 0 N_REFERENCE_GRID = 2000 # periodic spectral grid for the method-of-lines reference N_REFERENCE_STEPS = 2000 # classical RK4 time steps for that reference X_GRID = jnp.linspace(0.0, 1.0, N_REFERENCE_GRID, endpoint=False).reshape(-1, 1) def u0(x: jnp.ndarray) -> jnp.ndarray: """Initial condition, shared by both parts (kept in [0.2, 0.8]).""" return 0.5 + 0.3 * jnp.sin(2 * jnp.pi * x) def advection_field(t: jnp.ndarray, x: jnp.ndarray) -> jnp.ndarray: """Advection field a(t, x) = C (constant velocity).""" return jnp.ones_like(x) * C def exact_characteristic_foot(t: jnp.ndarray, x: jnp.ndarray, dt: float) -> jnp.ndarray: """Exact backward characteristic foot for `dx/ds = C`: `x - C * dt`.""" return x - C * dt def rate(x: jnp.ndarray) -> jnp.ndarray: """Spatially-varying logistic rate `r0 * (1 + 0.5 * sin(2 pi x))`.""" return R0 * (1 + 0.5 * jnp.sin(2 * jnp.pi * x)) def reaction_term(t: jnp.ndarray, x: jnp.ndarray, u: jnp.ndarray) -> jnp.ndarray: """Reaction term `s(t, x, u) = r(x) * u * (1 - u)`.""" return rate(x) * u * (1 - u) def f_init(x: jnp.ndarray) -> jnp.ndarray: return u0(x) def compose_splitting_steps(kind_reaction: str, dt: float, nt: int, n_substeps: int): """Compose `nt` Lie/Strang splitting steps analytically -- no network, no training. Mirrors `NeuralSemiLagrangian.solve`'s `target_func` exactly: reaction is applied per characteristic foot, before averaging (there is no diffusion here, so that average is trivial), with a second reaction half-step at the query point for "strang". `u^n` is represented as a plain Python closure composed over the previous step's closure, instead of a trained network. Returns: A callable `field(x) -> value at t = nt * dt`. """ characteristic = Characteristic( advection_field=advection_field, dt=dt, model_type="x", exact_characteristic_foot=exact_characteristic_foot, periodic=True, domain_bounds=DOMAIN_BOUNDS, ) reaction = ReactionStep(reaction_term, n_substeps=n_substeps) field = u0 for step in range(nt): t0 = step * dt prev_field = field def new_field(x, prev_field=prev_field, t0=t0): t0_arr = jnp.full_like(x, t0) feet = characteristic.compute_feet(t0_arr, x) def value_at_foot(foot): (foot_x,) = foot u_val = prev_field(foot_x) pre_tau = dt / 2 if kind_reaction == "strang" else dt return reaction.evolve(t0_arr, pre_tau, u_val, foot_x) target = sum(value_at_foot(foot) for foot in feet) / len(feet) if kind_reaction == "strang": t_mid = jnp.full_like(x, t0 + dt / 2) target = reaction.evolve(t_mid, dt / 2, target, x) return target field = new_field return field def fitted_order(dts: list[float], errors: list[float]) -> float: """Least-squares slope of log(error) vs log(dt): the observed order.""" slope, _ = np.polyfit(np.log(dts), np.log(errors), 1) return float(slope) def rk4_reference_solution( x_grid: jnp.ndarray, t_final: float, n_steps: int ) -> jnp.ndarray: """Independent reference for Part 2: method-of-lines + classical RK4 in time. Unlike Part 1's self-convergence check (which validates the Lie/Strang splitting scheme against a very fine instance of *itself*), this discretizes `du/dt = -C du/dx + r(x) u (1-u)` directly on a periodic grid -- the spatial derivative via a Fourier/spectral method (exact to machine precision here, for this smooth periodic profile, given `x_grid`'s resolution) -- and integrates the resulting ODE system with a plain RK4 step in time, with no operator splitting at all: a ground truth independent of NSL's own splitting machinery. Args: x_grid: Periodic grid (uniform, *not* including the right endpoint -- it is identified with the left one), shape (n, 1). t_final: Final time to integrate to (from t=0). n_steps: Number of (classical, unsplit) RK4 time steps. Returns: The reference solution at `t_final`, shape (n, 1). """ n = x_grid.shape[0] dx = float(x_grid[1, 0] - x_grid[0, 0]) wavenumbers = 2 * jnp.pi * jnp.fft.fftfreq(n, d=dx) rates = rate(x_grid).ravel() def dudt(u): du_dx = jnp.real(jnp.fft.ifft(1j * wavenumbers * jnp.fft.fft(u))) return -C * du_dx + rates * u * (1 - u) dt = t_final / n_steps def rk4_step(u, _): k1 = dudt(u) k2 = dudt(u + dt / 2 * k1) k3 = dudt(u + dt / 2 * k2) k4 = dudt(u + dt * k3) return u + dt / 6 * (k1 + 2 * k2 + 2 * k3 + k4), None u_final, _ = jax.lax.scan(rk4_step, u0(x_grid).ravel(), None, length=n_steps) return u_final.reshape(-1, 1) def make_trained_solver(kind: str, nt: int) -> NeuralSemiLagrangian: 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", reaction_term=reaction_term, kind_reaction=kind, n_reaction_substeps=4, ) if __name__ == "__main__": # ------------------------------------------------------------------ # Part 1: direct, untrained self-convergence order check. # ------------------------------------------------------------------ print("Part 1: local self-convergence order verification (no training)...") start = time.perf_counter() self_convergence_reference = compose_splitting_steps( "strang", TOTAL_TIME / PART1_REFERENCE_NT, PART1_REFERENCE_NT, PART1_N_REACTION_SUBSTEPS, )(X_PROBE) part1_errors: dict[str, list[float]] = {} part1_dts = [TOTAL_TIME / nt for nt in PART1_NTS] for kind in KINDS: errors = [] for nt in PART1_NTS: dt = TOTAL_TIME / nt field = compose_splitting_steps(kind, dt, nt, PART1_N_REACTION_SUBSTEPS) pred = field(X_PROBE) errors.append( float(jnp.sqrt(jnp.mean((pred - self_convergence_reference) ** 2))) ) part1_errors[kind] = errors order = fitted_order(part1_dts, errors) errs_str = ", ".join(f"{e:.2e}" for e in errors) print(f" {kind:>8s}: errors = [{errs_str}] (fitted order: {order:.2f})") print( " Theory: order 1 for 'lie', order 2 for 'strang' (self-convergence " f"against the nt={PART1_REFERENCE_NT} 'strang' reference). " f"({time.perf_counter() - start:.1f}s)\n" ) # ------------------------------------------------------------------ # Part 2: light, trained NeuralSemiLagrangian comparison. # ------------------------------------------------------------------ print("Part 2: light, trained NeuralSemiLagrangian comparison...") print(" Building the independent method-of-lines RK4 reference...") start = time.perf_counter() reference_field = rk4_reference_solution(X_GRID, TOTAL_TIME, N_REFERENCE_STEPS) reference_norm = jnp.sqrt(jnp.mean(reference_field**2)) print(f" ...done in {time.perf_counter() - start:.1f}s\n") key = jax.random.PRNGKey(SEED) nn = MLP(in_size=1, out_size=1, hidden_sizes=[16] * 2, 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, linesearch="armijo", adaptive_matrix_regularization=True, ) init_space = init_solver.space trained_errors: dict[str, list[float]] = {kind: [] for kind in KINDS} finest_profile: dict[str, jnp.ndarray] = {} 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, linesearch="armijo", adaptive_matrix_regularization=True, ) variables = neural_sl.space.create_variables() all_together = ParamVecFunction.cat(variables) batched_func = all_together.vmap_on_physical_variables() u_pred = batched_func(neural_sl.space, X_GRID) rel_error = float( jnp.sqrt(jnp.mean((u_pred - reference_field) ** 2)) / reference_norm ) trained_errors[kind].append(rel_error) if nt == NT_LIST_TRAINED[-1]: finest_profile[kind] = jax.device_get(u_pred).ravel() errs_str = ", ".join(f"{e:.2e}" for e in trained_errors[kind]) print(f" {kind:>8s} (nt={NT_LIST_TRAINED}): errors = [{errs_str}]") # ------------------------------------------------------------------ # Plot 1: Part 1's log-log self-convergence error, with order refs. # ------------------------------------------------------------------ fig1, ax1 = plt.subplots(figsize=(6, 5)) for kind in KINDS: ax1.loglog(part1_dts, part1_errors[kind], marker="o", label=kind) dt_ref = np.array([part1_dts[0], part1_dts[-1]]) for (order, style), kind in zip([(1, "k--"), (2, "k:")], KINDS): err_ref = part1_errors[kind][0] * (dt_ref / part1_dts[0]) ** order ax1.loglog(dt_ref, err_ref, style, label=f"order {order} (reference)") ax1.set_xlabel("dt") ax1.set_ylabel("self-convergence error (vs. fine 'strang' reference)") ax1.set_title("Splitting order (untrained)") ax1.legend() fig1.tight_layout() # ------------------------------------------------------------------ # Plot 2: Part 2's trained-solver error vs nt, against the RK4 reference. # ------------------------------------------------------------------ fig2, ax2 = plt.subplots(figsize=(6, 5)) dts_trained = jnp.array([(DOM_T[1] - DOM_T[0]) / nt for nt in NT_LIST_TRAINED]) for kind in KINDS: ax2.loglog(dts_trained, trained_errors[kind], marker="o", label=kind) for (order, style), kind in zip([(1, "k--"), (2, "k:")], KINDS): err_ref = trained_errors[kind][0] * (dts_trained / dts_trained[0]) ** order ax2.loglog(dts_trained, err_ref, style, label=f"order {order} (reference)") ax2.set_xlabel("dt") ax2.set_ylabel("final relative L2 error (vs. RK4 reference)") ax2.set_title("Trained NeuralSemiLagrangian (space-dependent rate)") ax2.legend() fig2.tight_layout() # ------------------------------------------------------------------ # Plot 3: final 1D profile (Part 2, nt=NT_LIST_TRAINED[-1]) vs the # RK4 reference. # ------------------------------------------------------------------ fig3, ax3 = plt.subplots(figsize=(7, 5)) x_grid_flat = jax.device_get(X_GRID).ravel() ax3.plot( x_grid_flat, jax.device_get(reference_field).ravel(), "k-", label="RK4 reference", linewidth=2, ) for kind in KINDS: ax3.plot(x_grid_flat, finest_profile[kind], "--", label=kind) ax3.set_xlabel("x") ax3.set_ylabel("u") ax3.set_title(f"Final profile at t={DOM_T[-1]} (nt={NT_LIST_TRAINED[-1]})") ax3.legend() fig3.tight_layout() plt.show()