r"""Validates ``NonCanonicalHamiltonianVectorFieldSpace`` and ``LagrangianDegenerateVectorFieldSpace`` (in ``vector_fields.py``) on the Lotka-Volterra predator-prey system, a textbook non-canonical Hamiltonian system -- no trainable network, no gradient, no Projector: both potentials are plain analytic (parametric) callables, plugged in exactly like an MLP flownet would be. 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). Part 1 -- NonCanonicalHamiltonianVectorFieldSpace: Lotka-Volterra's Poisson structure is W(q, p) = (1 / (q*p)) * [[0, 1], [-1, 0]], i.e. W = d(phi) - d(phi)^T for the 1-form potential phi(q, p) = (log(p) / q, 0). Part 2 -- LagrangianDegenerateVectorFieldSpace: its structure 1-form is fixed to (0, theta) -- a gauge-equivalent restriction of Part 1's phi (add d(chi) with chi(q, p) = -log(p)*log(q) to phi and its q-component cancels), giving theta(q, p) = -log(q) / p and the *same* W, hence the *same* flow. Both (phi, H) and (theta, H) below reproduce Lotka-Volterra exactly (derived by solving W @ F_LV = grad H for phi/theta by hand, then checked against an independent ``jax.jacobian``/``jax.grad`` reference). """ # %% 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 FlowsApproximationSpace, ) from scimba_jax.ode_approx.basic_discrete_ode_nets import ( Rk4Flow, ) from scimba_jax.ode_approx.vector_fields import ( LagrangianDegenerateVectorFieldSpace, NonCanonicalHamiltonianVectorFieldSpace, ) from scimba_jax.utils.scimba_pytree import ScimbaPytree jax.config.update("jax_enable_x64", True) # %% Hyperparameters a, b, c, d = 1.1, 0.4, 0.1, 0.4 # Lotka-Volterra params dt = 0.02 n_steps = 2000 q0, p0 = 2.0, 1.5 def lv_vector_field(u, mu): q, p = u[0], u[1] return jnp.array([q * (a - b * p), p * (c * q - d)]) def hamiltonian_lv(q, p): # noqa: N802 return c * q - d * jnp.log(q) + b * p - a * jnp.log(p) # %% Reference: hand-written RK4 integration of the true equations. def lv_rk4_step(x): def f(x): return lv_vector_field(x, None) k1 = f(x) k2 = f(x + dt / 2 * k1) k3 = f(x + dt / 2 * k2) k4 = f(x + dt * k3) return x + dt / 6 * (k1 + 2 * k2 + 2 * k3 + k4) def simulate_reference(x0, n_steps): def step_fn(carry, _): x_new = lv_rk4_step(carry) return x_new, x_new _, traj = jax.lax.scan(step_fn, x0, None, length=n_steps) return jnp.concatenate([x0[None], traj], axis=0) x0 = jnp.array([q0, p0]) traj_ref = simulate_reference(x0, n_steps) # %% Small ScimbaPytree wrapper for plain analytic callables (need `.ndof()` # to plug into ApproximationSpace like any other model). class AnalyticNet(ScimbaPytree): fn: object def __init__(self, fn): self.fn = fn def __call__(self, inputs): return self.fn(inputs) def ndof(self) -> int: return 0 def hamiltonian_exact(inputs): # noqa: N802 return jnp.array([hamiltonian_lv(inputs[0], inputs[1])]) # %% Part 1: NonCanonicalHamiltonianVectorFieldSpace with the exact (phi, H). def phi_exact(inputs): q, p = inputs[0], inputs[1] return jnp.array([jnp.log(p) / q, 0.0]) phi_space = ApproximationSpace( dims={"u": 2, "mu": 0}, list_models=[(AnalyticNet(phi_exact), "vec", 2)], model_type="u_mu", ) hamiltonian_space = ApproximationSpace( dims={"u": 2, "mu": 0}, list_models=[(AnalyticNet(hamiltonian_exact), "scalar", None)], model_type="u_mu", ) vf_phi = NonCanonicalHamiltonianVectorFieldSpace( dim=2, phi=phi_space, hamiltonian=hamiltonian_space ) rk4_phi = Rk4Flow(dim=2, flownet=vf_phi, dt=dt, time_dependent=False, params_dim=0) space_phi = FlowsApproximationSpace( state_dim=2, params_dim=0, model=rk4_phi, model_type="u_mu", rollout=1 ) traj_phi = space_phi.rollout_trajectory(space_phi, x0, jnp.array([]), n_steps) err_phi = jnp.max(jnp.abs(traj_phi - traj_ref)) print( "NonCanonicalHamiltonianVectorFieldSpace (phi, H): " f"max traj error vs reference = {float(err_phi):.4e}" ) # %% Part 2: LagrangianDegenerateVectorFieldSpace with the exact # (theta, H) -- theta = -log(q) / p, gauge-equivalent to phi above. def theta_exact(inputs): q, p = inputs[0], inputs[1] return jnp.array([-jnp.log(q) / p]) theta_space = ApproximationSpace( dims={"u": 2, "mu": 0}, list_models=[(AnalyticNet(theta_exact), "vec", 1)], model_type="u_mu", ) vf_theta = LagrangianDegenerateVectorFieldSpace( dim=2, theta=theta_space, hamiltonian=hamiltonian_space ) rk4_theta = Rk4Flow(dim=2, flownet=vf_theta, dt=dt, time_dependent=False, params_dim=0) space_theta = FlowsApproximationSpace( state_dim=2, params_dim=0, model=rk4_theta, model_type="u_mu", rollout=1 ) traj_theta = space_theta.rollout_trajectory(space_theta, x0, jnp.array([]), n_steps) err_theta = jnp.max(jnp.abs(traj_theta - traj_ref)) print( "LagrangianDegenerateVectorFieldSpace (theta, H): " f"max traj error vs reference = {float(err_theta):.4e}" ) # %% Plots t_axis = jnp.linspace(0, n_steps * dt, n_steps + 1) fig, axes = plt.subplots(1, 3, figsize=(16, 5)) ax = axes[0] ax.plot(traj_ref[:, 0], traj_ref[:, 1], "k-", label="reference RK4", linewidth=2) ax.plot(traj_phi[:, 0], traj_phi[:, 1], "C0--", label="NonCanonical (phi, H)") ax.plot( traj_theta[:, 0], traj_theta[:, 1], "C1--", label="LagrangianDegenerate (theta, H)" ) ax.set_xlabel("q (prey)") ax.set_ylabel("p (predator)") ax.set_title("Phase portrait") ax.legend() ax = axes[1] ax.plot(t_axis, traj_ref[:, 0], "k-", label="reference RK4", linewidth=2) ax.plot(t_axis, traj_phi[:, 0], "C0--", label="NonCanonical (phi, H)") ax.plot(t_axis, traj_theta[:, 0], "C1--", label="LagrangianDegenerate (theta, H)") ax.set_xlabel("t") ax.set_ylabel("q(t)") ax.set_title("Prey population over time") ax.legend() ax = axes[2] H_ref = jax.vmap(lambda xy: hamiltonian_lv(xy[0], xy[1]))(traj_ref) H_phi = jax.vmap(lambda xy: hamiltonian_lv(xy[0], xy[1]))(traj_phi) H_theta = jax.vmap(lambda xy: hamiltonian_lv(xy[0], xy[1]))(traj_theta) ax.plot(t_axis, H_ref - H_ref[0], "k-", label="reference RK4", linewidth=2) ax.plot(t_axis, H_phi - H_phi[0], "C0--", label="NonCanonical (phi, H)") ax.plot(t_axis, H_theta - H_theta[0], "C1--", label="LagrangianDegenerate (theta, H)") ax.set_xlabel("t") ax.set_ylabel("H(t) - H(0)") ax.set_title("Conservation of H") ax.legend() plt.tight_layout() plt.show() # %%