r"""Validates ``SymplecticEulerFlowNonSep`` and ``VerletFlowNonSep`` (the non-separable generalizations of ``SymplecticEulerFlowSep``/ ``VerletFlowSep``, in ``symplec_discrete_ode_nets.py``) against ``Rk4Flow`` on a genuinely non-separable Hamiltonian -- no trainable network, no gradient, no Projector (see ``pendulum_batch_integration.py``). Dynamics: a pendulum-like system with position-dependent mass, H(q, p) = 0.5 * p^2 * (1 + q^2) + 0.5 * q^2 -- the kinetic term couples q and p, so H cannot be split into K(p) + U(q): both non-separable flow classes need an implicit (Newton) solve internally, unlike their separable counterparts. RK4 is 4th order (much more accurate locally) but *not* symplectic: its energy error grows without bound over long integrations. The two symplectic schemes are only 1st/2nd order (larger local error) but their energy error stays *bounded* forever -- the classic trade-off between local accuracy and long-time qualitative correctness that motivates symplectic integrators in the first place. """ # %% 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, PhaseSpaceFlowApproximationSpace, ) from scimba_jax.ode_approx.basic_discrete_ode_nets import ( Rk4Flow, ) from scimba_jax.ode_approx.symplec_discrete_ode_nets import ( SymplecticEulerFlowNonSep, VerletFlowNonSep, ) jax.config.update("jax_enable_x64", True) # %% Hyperparameters dt = 0.07 n_steps = 40000 # Tf = 1000 newton_iter = 20 # %% Non-separable Hamiltonian: position-dependent mass. Analytic callables, # no trainable state. def hamiltonian(q, p, mu): return jnp.array([0.5 * p[0] ** 2 * (1.0 + q[0] ** 2) + 0.5 * q[0] ** 2]) def hamiltonian_vector_field(u, mu): """[dq/dt, dp/dt] = [dH/dp, -dH/dq], u = [q, p] merged (Rk4Flow's "u_mu" convention).""" q, p = u[0], u[1] dH_dq = p**2 * q + q dH_dp = p * (1.0 + q**2) return jnp.array([dH_dp, -dH_dq]) def energy(q, p): return 0.5 * p**2 * (1.0 + q**2) + 0.5 * q**2 # %% Batch of initial conditions N_ic = 4 q0s = jnp.array([0.5, 1.0, 1.5, 2.0]) p0s = jnp.array([0.5, 0.5, 0.3, 0.2]) # %% SymplecticEulerFlowNonSep (1st order, symplectic) symp_model = SymplecticEulerFlowNonSep( dim=2, hamiltonian_net=hamiltonian, dt=dt, time_dependent=False, params_dim=0, newton_iter=newton_iter, ) space_symp = PhaseSpaceFlowApproximationSpace( state_dim=1, params_dim=0, model=symp_model, model_type="q_p_mu", rollout=1, ) def integrate_symp(q0, p0): return space_symp.rollout_trajectory( space_symp, jnp.array([q0]), jnp.array([p0]), jnp.array([]), n_steps ) q_symp, p_symp = jax.vmap(integrate_symp)(q0s, p0s) # each (N_ic, n_steps + 1, 1) # %% VerletFlowNonSep (2nd order, symplectic) verlet_model = VerletFlowNonSep( dim=2, hamiltonian_net=hamiltonian, dt=dt, time_dependent=False, params_dim=0, newton_iter=newton_iter, ) space_verlet = PhaseSpaceFlowApproximationSpace( state_dim=1, params_dim=0, model=verlet_model, model_type="q_p_mu", rollout=1, ) def integrate_verlet(q0, p0): return space_verlet.rollout_trajectory( space_verlet, jnp.array([q0]), jnp.array([p0]), jnp.array([]), n_steps ) q_verlet, p_verlet = jax.vmap(integrate_verlet)(q0s, p0s) # %% Rk4Flow (4th order, not symplectic) rk4_model = Rk4Flow( dim=2, flownet=hamiltonian_vector_field, dt=dt, time_dependent=False, params_dim=0 ) space_rk4 = FlowsApproximationSpace( state_dim=2, params_dim=0, model=rk4_model, model_type="u_mu", rollout=1 ) def integrate_rk4(q0, p0): x0 = jnp.array([q0, p0]) return space_rk4.rollout_trajectory(space_rk4, x0, jnp.array([]), n_steps) traj_rk4 = jax.vmap(integrate_rk4)(q0s, p0s) # (N_ic, n_steps + 1, 2) # %% Energy drift: symplectic schemes should stay bounded, RK4 should grow. H_symp = jax.vmap(energy)(q_symp[..., 0], p_symp[..., 0]) H_verlet = jax.vmap(energy)(q_verlet[..., 0], p_verlet[..., 0]) H_rk4 = jax.vmap(energy)(traj_rk4[..., 0], traj_rk4[..., 1]) half = n_steps // 2 print(f"{'scheme':>22} {'drift 1st half':>16} {'drift 2nd half':>16}") for name, H_traj in [ ("SymplecticEulerNonSep", H_symp), ("VerletNonSep", H_verlet), ("Rk4", H_rk4), ]: d1 = float(jnp.max(jnp.abs(H_traj[:, :half] - H_traj[:, :1]))) d2 = float(jnp.max(jnp.abs(H_traj[:, half:] - H_traj[:, :1]))) print(f"{name:>22} {d1:16.4e} {d2:16.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] for i in range(N_ic): ax.plot(q_symp[i, :2000, 0], p_symp[i, :2000, 0], "C0-", alpha=0.6) ax.plot(q_verlet[i, :2000, 0], p_verlet[i, :2000, 0], "C1-", alpha=0.6) ax.plot(traj_rk4[i, :2000, 0], traj_rk4[i, :2000, 1], "C2-", alpha=0.6) ax.plot([], [], "C0-", label="SymplecticEulerNonSep") ax.plot([], [], "C1-", label="VerletNonSep") ax.plot([], [], "C2-", label="Rk4") ax.set_xlabel("q") ax.set_ylabel("p") ax.set_title("Phase portraits (t in [0, 100])") ax.legend() ax = axes[1] for i in range(N_ic): ax.plot(t_axis, H_symp[i] - H_symp[i, 0], "C0-", alpha=0.6) ax.plot(t_axis, H_verlet[i] - H_verlet[i, 0], "C1-", alpha=0.6) ax.plot(t_axis, H_rk4[i] - H_rk4[i, 0], "C2-", alpha=0.6) ax.plot([], [], "C0-", label="SymplecticEulerNonSep") ax.plot([], [], "C1-", label="VerletNonSep") ax.plot([], [], "C2-", label="Rk4") ax.set_xlabel("t") ax.set_ylabel("H(t) - H(0)") ax.set_title("Energy drift over the full run (t in [0, 1000])") ax.legend() ax = axes[2] running_max_symp = jax.vmap(lambda h: jax.lax.cummax(jnp.abs(h - h[0])))(H_symp) running_max_verlet = jax.vmap(lambda h: jax.lax.cummax(jnp.abs(h - h[0])))(H_verlet) running_max_rk4 = jax.vmap(lambda h: jax.lax.cummax(jnp.abs(h - h[0])))(H_rk4) for i in range(N_ic): ax.semilogy(t_axis, running_max_symp[i], "C0-", alpha=0.6) ax.semilogy(t_axis, running_max_verlet[i], "C1-", alpha=0.6) ax.semilogy(t_axis, running_max_rk4[i], "C2-", alpha=0.6) ax.plot([], [], "C0-", label="SymplecticEulerNonSep") ax.plot([], [], "C1-", label="VerletNonSep") ax.plot([], [], "C2-", label="Rk4") ax.set_xlabel("t") ax.set_ylabel("running max |H(t) - H(0)|") ax.set_title("Bounded (symplectic) vs growing (Rk4) energy error") ax.legend() plt.tight_layout() plt.show() # %%