r"""Validates the explicit Adams-Bashforth multistep flows (AB2, AB3) as plain numerical integrators -- no trainable network, no gradient, no Projector (see ``implicit_vs_explicit_euler.py``). Dynamics (same cubic decay as ``implicit_vs_explicit_euler.py``): dx/dt = -mu * x^3, x(0) = x0 > 0 with closed-form solution x(t) = x0 / sqrt(1 + 2*mu*x0^2*t). AB2/AB3 aren't self-starting: a k-step method needs the k most recent states before it can take its first step, which a single x0 doesn't give. ``bootstrap_multistep_window`` (in ``basic_discrete_ode_nets.py``) takes k-1 RK4 steps from x0 to build the initial window (x_{k-1}, ..., x_0) that ``MultistepFlowsApproxSpace.rollout_trajectory`` expects. The scheme itself (``AdamsBashforth{2,3}Flow.step``) is written purely with ParamFunc algebra over NAMED window variables ("un", "unm1", "unm2") via ``arg``/``pullback`` -- no array slicing; ``MultistepFlowsApproxSpace`` is the one that shifts the window between steps. This script checks the expected convergence order (max error over [0, Tf] as a function of dt, log-log slope): 2 for AB2, 3 for AB3, against RK4 (order 4) and Explicit Euler (order 1) as references. """ # %% import jax import jax.numpy as jnp import matplotlib.pyplot as plt from scimba_jax.nonlinear_approximation.approximation_spaces.flow_approximation_spaces import ( # noqa: E501 FlowsApproximationSpace, MultistepFlowsApproxSpace, ) from scimba_jax.ode_approx.basic_discrete_ode_nets import ( AdamsBashforth2Flow, AdamsBashforth3Flow, ExplicitEulerFlow, Rk4Flow, bootstrap_multistep_window, ) jax.config.update("jax_enable_x64", True) # %% Cubic-damping dynamics -- analytic callable, no trainable state. DIM = 1 def cubic_vector_field(x, mu): return jnp.array([-mu[0] * x[0] ** 3]) def cubic_analytic(t, x0, mu): return x0 / jnp.sqrt(1.0 + 2.0 * mu * x0**2 * t) x0_val = 2.0 mu_val = 1.0 x0 = jnp.array([x0_val]) mu = jnp.array([mu_val]) Tf = 2.0 # model_type per window depth k (2 -> AB2, 3 -> AB3) _MULTISTEP_MODEL_TYPE = {2: "un_unm1_mu", 3: "un_unm1_unm2_mu"} # %% Integrate a fixed physical time span Tf with a given (flow_class, k). # # k = 1 for the one-step references (Euler, RK4): plain FlowsApproximationSpace # rollout from x0. k > 1 for the multistep flows: bootstrap the initial window # with an RK4 starter, then roll out the remaining steps on a # MultistepFlowsApproxSpace; the trajectory is always the physical "un" (most # recent) component of the window. def integrate(flow_class, k, dt): n_steps = int(round(Tf / dt)) model = flow_class(dim=DIM, flownet=cubic_vector_field, dt=dt, params_dim=1) if k == 1: space = FlowsApproximationSpace( state_dim=DIM, params_dim=1, model=model, model_type="u_mu", rollout=1 ) traj = space.rollout_trajectory(space, x0, mu, n_steps) t_axis = jnp.arange(n_steps + 1) * dt else: space = MultistepFlowsApproxSpace( dim=DIM, params_dim=1, model=model, model_type=_MULTISTEP_MODEL_TYPE[k] ) starter = Rk4Flow(dim=DIM, flownet=cubic_vector_field, dt=dt, params_dim=1) window0 = bootstrap_multistep_window(starter, x0, mu, k) n_multistep = n_steps - (k - 1) traj = space.rollout_trajectory(space, window0, mu, n_multistep) t_axis = (jnp.arange(n_multistep + 1) + (k - 1)) * dt return t_axis, traj[:, 0] methods = { "Euler (order 1)": (ExplicitEulerFlow, 1), "RK4 (order 4)": (Rk4Flow, 1), "AB2 (order 2)": (AdamsBashforth2Flow, 2), "AB3 (order 3)": (AdamsBashforth3Flow, 3), } # %% Convergence-order study: max error over [0, Tf] vs dt # The AB3/RK4 asymptotic order only kicks in once dt is small enough that # the leading truncation-error term dominates: dt=0.2 is too coarse for a # cubic nonlinearity this strong (mu*x0^2 dt ~ 1) and would show apparent # sub-order behavior, so the sweep starts at 0.1. dts = jnp.array([0.1, 0.05, 0.025, 0.0125, 0.00625]) errors = {name: [] for name in methods} for dt in dts: dt = float(dt) for name, (flow_class, k) in methods.items(): t_axis, x_traj = integrate(flow_class, k, dt) exact = cubic_analytic(t_axis, x0_val, mu_val) errors[name].append(float(jnp.max(jnp.abs(x_traj - exact)))) print(f"{'dt':>8}" + "".join(f"{name:>18}" for name in methods)) for i, dt in enumerate(dts): row = "".join(f"{errors[name][i]:18.3e}" for name in methods) print(f"{float(dt):8.4f}{row}") print( "\nObserved order (slope of log(err) vs log(dt), consecutive dt pairs)." "\nRK4's order goes erratic at the smallest dt: its error has already" "\nhit the float64 round-off floor (~1e-9), so its slope there measures" "\nnoise, not truncation error." ) for name in methods: errs = jnp.array(errors[name]) orders = jnp.log(errs[:-1] / errs[1:]) / jnp.log(dts[:-1] / dts[1:]) print(f" {name:20s}: {[f'{float(o):.2f}' for o in orders]}") # %% Trajectory plot at a representative dt, + convergence-order plot dt_plot = 0.1 fig, axes = plt.subplots(1, 2, figsize=(12, 5)) t_fine = jnp.linspace(0, Tf, 400) axes[0].plot(t_fine, cubic_analytic(t_fine, x0_val, mu_val), "k-", label="analytic") for name, (flow_class, k) in methods.items(): t_axis, x_traj = integrate(flow_class, k, dt_plot) axes[0].plot(t_axis, x_traj, "--", label=name) axes[0].set_xlabel("t") axes[0].set_ylabel("x(t)") axes[0].set_title(f"Trajectories, dt={dt_plot}") axes[0].legend() for name in methods: axes[1].loglog(dts, errors[name], "o-", label=name) axes[1].set_xlabel("dt") axes[1].set_ylabel("max error over [0, Tf]") axes[1].set_title("Convergence order") axes[1].legend() axes[1].invert_xaxis() plt.tight_layout() plt.show() # %%