r"""Integrates the pendulum ODE for a batch of initial conditions, using FlowsApproximationSpace / PhaseSpaceFlowApproximationSpace as plain numerical integrators -- no trainable network, no gradient, no Projector. ``Rk2Flow``, ``SymplecticEulerFlowSep`` and ``VerletFlowSep`` accept a bare analytic callable as their vector field / potentials (see ``basic_discrete_ode_nets.py`` / ``symplec_discrete_ode_nets.py``), so the whole pipeline here is just a batched (``jax.vmap``) classical ODE solver built out of the flow algebra. Dynamics given by the Hamiltonian: H(q, p) = p^2/2 + mu*q^2/2 + mu*0.012*q^3/3 """ # %% 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 ( Rk2Flow, ) from scimba_jax.ode_approx.symplec_discrete_ode_nets import ( SymplecticEulerFlowSep, VerletFlowSep, ) jax.config.update("jax_enable_x64", True) # %% Hyperparameters dt = 0.1 Tf = 200.0 n_steps = int(Tf / dt) mu = jnp.array([1.0]) # %% Pendulum dynamics -- analytic callables, no trainable state. def pendulum_vector_field(u, mu): """Hamiltonian vector field [dq/dt, dp/dt], u = [q, p] merged (Rk2Flow's "u_mu" convention).""" q, p = u[0], u[1] m = mu[0] dq = p dp = -m * q - m * 0.012 * q**2 return jnp.array([dq, dp]) def pendulum_potentials(q_arr, p_arr, mu): """[K, U] with K(p) = p^2/2, U(q) = mu*q^2/2 + mu*0.012*q^3/3 (SymplecticEulerFlowSep's "q_p_mu" convention).""" q = q_arr[0] p = p_arr[0] m = mu[0] K = 0.5 * p**2 U = 0.5 * m * q**2 + 0.012 * m * q**3 / 3 return jnp.array([K, U]) # %% Batch of initial conditions (same mu, different (q0, p0)) N_ic = 8 q0s = jnp.linspace(1.0, 3.0, N_ic) p0s = jnp.zeros(N_ic) # %% Case 1: Rk2Flow (generic RK2 integrator, not structure-preserving) rk2_model = Rk2Flow( dim=2, flownet=pendulum_vector_field, dt=dt, time_dependent=False, params_dim=1 ) space_rk2 = FlowsApproximationSpace( state_dim=2, params_dim=1, model=rk2_model, model_type="u_mu", rollout=1 ) def integrate_rk2(q0, p0): x0 = jnp.array([q0, p0]) return space_rk2.rollout_trajectory(space_rk2, x0, mu, n_steps) traj_rk2 = jax.vmap(integrate_rk2)(q0s, p0s) # (N_ic, n_steps + 1, 2) # %% Case 2: SymplecticEulerFlowSep (structure-preserving) symp_model = SymplecticEulerFlowSep( dim=2, potentials_net=pendulum_potentials, dt=dt, time_dependent=False, params_dim=1, ) space_symp = PhaseSpaceFlowApproximationSpace( state_dim=1, params_dim=1, model=symp_model, model_type="q_p_mu", rollout=1, ) def integrate_symp(q0, p0): x0 = jnp.array([q0]) v0 = jnp.array([p0]) return space_symp.rollout_trajectory(space_symp, x0, v0, mu, n_steps) q_traj, p_traj = jax.vmap(integrate_symp)(q0s, p0s) # each (N_ic, n_steps + 1, 1) # %% Case 3: VerletFlowSep (Störmer-Verlet, 2nd-order structure-preserving) verlet_model = VerletFlowSep( dim=2, potentials_net=pendulum_potentials, dt=dt, time_dependent=False, params_dim=1, ) space_verlet = PhaseSpaceFlowApproximationSpace( state_dim=1, params_dim=1, model=verlet_model, model_type="q_p_mu", rollout=1, ) def integrate_verlet(q0, p0): x0 = jnp.array([q0]) v0 = jnp.array([p0]) return space_verlet.rollout_trajectory(space_verlet, x0, v0, mu, n_steps) q_traj_verlet, p_traj_verlet = jax.vmap(integrate_verlet)(q0s, p0s) # %% Energy drift: the symplectic schemes should stay bounded, RK2 should drift. def hamiltonian(q, p, m): return 0.5 * p**2 + 0.5 * m * q**2 + 0.012 * m * q**3 / 3 H_rk2 = hamiltonian(traj_rk2[..., 0], traj_rk2[..., 1], mu[0]) # (N_ic, n_steps + 1) H_symp = hamiltonian(q_traj[..., 0], p_traj[..., 0], mu[0]) H_verlet = hamiltonian(q_traj_verlet[..., 0], p_traj_verlet[..., 0], mu[0]) # %% Plots fig, axes = plt.subplots(1, 4, figsize=(20, 5)) ax = axes[0] for i in range(N_ic): ax.plot(traj_rk2[i, :, 0], traj_rk2[i, :, 1]) ax.set_xlabel("q") ax.set_ylabel("p") ax.set_title("Rk2Flow phase portrait") ax = axes[1] for i in range(N_ic): ax.plot(q_traj[i, :, 0], p_traj[i, :, 0]) ax.set_xlabel("q") ax.set_ylabel("p") ax.set_title("SymplecticEulerFlowSep phase portrait") ax = axes[2] for i in range(N_ic): ax.plot(q_traj_verlet[i, :, 0], p_traj_verlet[i, :, 0]) ax.set_xlabel("q") ax.set_ylabel("p") ax.set_title("VerletFlowSep phase portrait") ax = axes[3] t_axis = jnp.linspace(0, n_steps * dt, n_steps + 1) for i in range(N_ic): ax.plot(t_axis, H_rk2[i] - H_rk2[i, 0], "C0-", alpha=0.5) ax.plot(t_axis, H_symp[i] - H_symp[i, 0], "C1-", alpha=0.5) ax.plot(t_axis, H_verlet[i] - H_verlet[i, 0], "C2-", alpha=0.5) ax.plot([], [], "C0-", label="Rk2Flow") ax.plot([], [], "C1-", label="SymplecticEulerFlowSep") ax.plot([], [], "C2-", label="VerletFlowSep") ax.set_xlabel("t") ax.set_ylabel("H(t) - H(0)") ax.set_title("Energy drift") ax.legend() plt.tight_layout() plt.show() # %%