r"""Validates ``DegenerateVariationalIntegratorFlow`` (DVI) on the Lotka-Volterra predator-prey system, with the CORRECT structure 1-form theta(q, p) = -log(p) / q. Dynamics: dq/dt = q * (a - b*p) (q = prey) dp/dt = p * (c*q - d) (p = predator) with conserved quantity H(q, p) = c*q - d*log(q) + b*p - a*log(p). An earlier version of this script used theta(q, p) = -log(q) / p (q and p swapped inside the formula) and needed an ad hoc stabilizing gauge shift (theta -> theta + K/p) to avoid NaN divergence, yet still only reached order-1 local accuracy with a large error constant. Cross-checking against the reference implementation of Franck et al. (github.com/tremelow/ symplearn -- `experiments/lotka_volterra/models.py`, `src/symplearn/dynamics.py`) revealed q and p were swapped inside theta: their ``oneform(x, y) = -log(y) / x`` is, with x = prey = q and y = predator = p (same roles as here), theta(q, p) = -log(p) / q -- the OPPOSITE of what this script originally had. With the correct orientation, the DVI needs no gauge shift at all: it is directly stable and reaches order 2 (not 1) local truncation accuracy -- the swapped version wasn't merely a worse gauge with a larger error constant, it broke the scheme's convergence order outright. Caveat for anyone reusing this derivation elsewhere in this codebase: ``LagrangianDegenerateVectorFieldSpace`` (in ``vector_fields.py``, not touched by this script) computes the continuous flow using D_q(theta) for BOTH qdot and pdot; the reference implementation's general ``AbstractDegenLagrangian.vector_field`` uses D_p(theta) for both instead. Whether this is an actual bug in ``LagrangianDegenerateVectorFieldSpace`` or an equally-valid alternative convention has NOT been resolved here -- this script deliberately avoids it entirely, bootstrapping and validating against a plain hand-written RK4 integration of the TRUE (q, p) equations above instead of going through that class. """ # %% import jax import jax.numpy as jnp import matplotlib.pyplot as plt from scimba_jax.nonlinear_approximation.approximation_spaces.approximation_spaces import ( # noqa: E501 ApproximationSpace, ) from scimba_jax.nonlinear_approximation.approximation_spaces.flow_approximation_spaces import ( # noqa: E501 MultistepPhaseSpaceApproxSpace, ) from scimba_jax.ode_approx.symplec_discrete_ode_nets import ( DegenerateVariationalIntegratorFlow, ) from scimba_jax.utils.scimba_pytree import ScimbaPytree jax.config.update("jax_enable_x64", True) # %% Lotka-Volterra, theta(q, p) = -log(p) / q a, b, c, d = 1.1, 0.4, 0.1, 0.4 # Lotka-Volterra params q0, p0 = 2.0, 1.5 class AnalyticNet(ScimbaPytree): """Plain analytic callable wrapped as a (parameter-free) ScimbaPytree model.""" fn: object def __init__(self, fn): self.fn = fn def __call__(self, inputs): return self.fn(inputs) def ndof(self) -> int: return 0 def theta_fn(u): q, p = u[0:1], u[1:2] return -jnp.log(p) / q def hamiltonian_fn(u): q, p = u[0:1], u[1:2] return c * q - d * jnp.log(q) + b * p - a * jnp.log(p) def hamiltonian(q, p): return float(c * q - d * jnp.log(q) + b * p - a * jnp.log(p)) potentials_qp = ApproximationSpace( dims={"q": 1, "p": 1, "mu": 0}, list_models=[ (AnalyticNet(theta_fn), "vec", 1), (AnalyticNet(hamiltonian_fn), "scalar", None), ], model_type="q_p_mu", ) u0 = jnp.array([q0, p0]) mu = jnp.array([]) # %% Reference: hand-written RK4 integration of the true (continuous) LV # equations -- also used as the (very fine dt) ground truth for the local # truncation order test below, and to bootstrap the DVI's 2-point window # (deliberately NOT using LagrangianDegenerateVectorFieldSpace, see the # module docstring). def lv_vector_field(u): q, p = u[0], u[1] return jnp.array([q * (a - b * p), p * (c * q - d)]) def lv_rk4_step(x, dt): k1 = lv_vector_field(x) k2 = lv_vector_field(x + dt / 2 * k1) k3 = lv_vector_field(x + dt / 2 * k2) k4 = lv_vector_field(x + dt * k3) return x + dt / 6 * (k1 + 2 * k2 + 2 * k3 + k4) def lv_exact(x0, t, n_fine=20000): dt_fine = t / n_fine def step_fn(carry, _): x_new = lv_rk4_step(carry, dt_fine) return x_new, None x_final, _ = jax.lax.scan(step_fn, x0, None, length=n_fine) return x_final def rollout_dvi(dt: float, total_time: float): n_steps = int(total_time / dt) dvi_model = DegenerateVariationalIntegratorFlow( dim=2, potentials_hamiltonian_space=potentials_qp, dt=dt, newton_iter=25 ) dvi_space = MultistepPhaseSpaceApproxSpace( state_dim=1, params_dim=0, model=dvi_model, model_type="q_p_qm_pm_mu" ) x1 = lv_rk4_step(u0, dt) # bootstrap: one true RK4 step, not via a Flow class window0 = (x1[:1], x1[1:], u0[:1], u0[1:]) return dvi_space.rollout_trajectory(dvi_space, window0, mu, n_steps) # %% Local truncation order: one DVI step (from an RK4-bootstrapped window) # vs a very fine RK4 reference, as dt shrinks -- order 2 (not 1) with the # correct theta orientation. dts = [0.02, 0.01, 0.005, 0.0025, 0.00125] errors = [] for dt in dts: qs, ps = rollout_dvi(dt, total_time=2 * dt) u_dvi = jnp.array([qs[1, 0], ps[1, 0]]) u_exact = lv_exact(u0, 2 * dt) errors.append(float(jnp.max(jnp.abs(u_dvi - u_exact)))) print(f"{'dt':>8}{'error':>14}") for dt, e in zip(dts, errors): print(f"{dt:8.5f}{e:14.4e}") orders = [ jnp.log(errors[i] / errors[i + 1]) / jnp.log(dts[i] / dts[i + 1]) for i in range(len(dts) - 1) ] print( "Observed order (consecutive dt pairs, expect ~2):", [f"{float(o):.2f}" for o in orders], ) # %% Long-time behavior: stable over many oscillation periods, tracking the # continuous reference closely -- no gauge shift needed. dt_long = 0.005 total_time = 40.0 qs_dvi, ps_dvi = rollout_dvi(dt_long, total_time) n_steps = qs_dvi.shape[0] - 1 diverged = bool(jnp.any(jnp.isnan(qs_dvi)) or jnp.any(jnp.isnan(ps_dvi))) print(f"\nDiverges (NaN) = {diverged}") H0 = hamiltonian(q0, p0) H_dvi = jnp.array( [hamiltonian(float(q), float(p)) for q, p in zip(qs_dvi[:, 0], ps_dvi[:, 0])] ) print( f"q range=[{float(qs_dvi.min()):.3f}, {float(qs_dvi.max()):.3f}] " f"p range=[{float(ps_dvi.min()):.3f}, {float(ps_dvi.max()):.3f}] " f"H rel. range={(float(H_dvi.max()) - float(H_dvi.min())) / H0:.3e}" ) # qs_dvi[0]/ps_dvi[0] is window0[0] = the bootstrapped state at t=dt_long # (see MultistepPhaseSpaceApproxSpace.rollout_trajectory's docstring), not # u0 at t=0 -- start the reference trajectory from that SAME time. x1_ref = lv_rk4_step(u0, dt_long) def step_fn(carry, _): x_new = lv_rk4_step(carry, dt_long) return x_new, x_new _, traj_ref = jax.lax.scan(step_fn, x1_ref, None, length=n_steps) traj_ref = jnp.concatenate([x1_ref[None], traj_ref], axis=0) rms_q = float(jnp.sqrt(jnp.mean((qs_dvi[:, 0] - traj_ref[:, 0]) ** 2))) print(f"RMS q error over {total_time:.0f} t.u. = {rms_q:.4f}") # %% Plots t_axis = dt_long + jnp.arange(n_steps + 1) * dt_long fig, axes = plt.subplots(1, 4, figsize=(20, 5)) axes[0].loglog(dts, errors, "o-") axes[0].set_xlabel("dt") axes[0].set_ylabel("error after 1 DVI step") axes[0].set_title("Local truncation order (~2)") ax = axes[1] ax.plot(traj_ref[:, 0], traj_ref[:, 1], "k-", label="reference RK4", linewidth=2) ax.plot(qs_dvi[:, 0], ps_dvi[:, 0], "C0--", label="DVI") ax.set_xlabel("q (prey)") ax.set_ylabel("p (predator)") ax.set_title("Phase portrait") ax.legend() ax = axes[2] ax.plot(t_axis, traj_ref[:, 0], "k-", label="reference RK4", linewidth=2) ax.plot(t_axis, qs_dvi[:, 0], "C0--", label="DVI") ax.set_xlabel("t") ax.set_ylabel("q(t)") ax.set_title("Prey population over time") ax.legend() ax = axes[3] ax.plot(t_axis, H_dvi - H0, "C0-") ax.set_xlabel("t") ax.set_ylabel("H(t) - H0") ax.set_title("Energy conservation") plt.tight_layout() plt.show() # %%