r"""Integrates a weakly nonlinear (cubic damping) ODE for a batch of decay rates ``mu``, using ``ExplicitEulerFlow`` and ``ImplicitEulerFlow`` as plain numerical integrators -- no trainable network, no gradient, no Projector (see ``pendulum_batch_integration.py``/``sir_batch_integration.py``). Dynamics: dx/dt = -mu * x^3, x(0) = x0 > 0 with closed-form solution x(t) = x0 / sqrt(1 + 2*mu*x0^2*t), monotonically decaying to 0. This is a classic textbook example of explicit-Euler instability on a dissipative nonlinearity: for large mu*dt*x0^2, the first explicit step overshoots past 0 into a much larger negative value (the cubic RHS grows fast), and the scheme oscillates with growing amplitude instead of decaying. Implicit Euler's per-step equation x + dt*mu*x^3 = x_old is strictly increasing in x (unique real root for any dt), so Newton converges reliably from x0 = x_old and the scheme stays bounded and monotonically decaying for any step size (unconditional stability). """ # %% 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, ) from scimba_jax.ode_approx.basic_discrete_ode_nets import ( ExplicitEulerFlow, ImplicitEulerFlow, ) jax.config.update("jax_enable_x64", True) # %% Hyperparameters dt = 0.1 n_steps = 50 x0 = 2.0 # %% Cubic-damping dynamics -- analytic callable, no trainable state. def cubic_vector_field(u, mu): return jnp.array([-mu[0] * u[0] ** 3]) def cubic_analytic(t, x0, mu): return x0 / jnp.sqrt(1.0 + 2.0 * mu * x0**2 * t) # %% Batch of decay rates mu mus = jnp.array([0.5, 1.0, 2.0, 3.0, 4.0, 5.0]) N_mu = mus.shape[0] # %% Explicit Euler explicit_model = ExplicitEulerFlow( dim=1, flownet=cubic_vector_field, dt=dt, time_dependent=False, params_dim=1 ) space_explicit = FlowsApproximationSpace( state_dim=1, params_dim=1, model=explicit_model, model_type="u_mu", rollout=1 ) def integrate_explicit(mu): return space_explicit.rollout_trajectory( space_explicit, jnp.array([x0]), jnp.array([mu]), n_steps ) traj_explicit = jax.vmap(integrate_explicit)(mus) # (N_mu, n_steps + 1, 1) # %% Implicit Euler implicit_model = ImplicitEulerFlow( dim=1, flownet=cubic_vector_field, dt=dt, time_dependent=False, params_dim=1, newton_iter=15, ) space_implicit = FlowsApproximationSpace( state_dim=1, params_dim=1, model=implicit_model, model_type="u_mu", rollout=1 ) def integrate_implicit(mu): return space_implicit.rollout_trajectory( space_implicit, jnp.array([x0]), jnp.array([mu]), n_steps ) traj_implicit = jax.vmap(integrate_implicit)(mus) # (N_mu, n_steps + 1, 1) # %% Reference: analytic solution t_axis = jnp.linspace(0, n_steps * dt, n_steps + 1) traj_analytic = jax.vmap(lambda mu: cubic_analytic(t_axis, x0, mu))(mus) print( f"{'mu':>5} {'implicit stays in [0,x0]':>26} {'explicit stays in [0,x0]':>26} {'err_implicit':>14} {'err_explicit':>14}" ) for i in range(N_mu): imp_bounded = bool( jnp.all((traj_implicit[i, :, 0] >= 0) & (traj_implicit[i, :, 0] <= x0)) ) exp_bounded = bool( jnp.all((traj_explicit[i, :, 0] >= 0) & (traj_explicit[i, :, 0] <= x0)) ) err_i = float(jnp.max(jnp.abs(traj_implicit[i, :, 0] - traj_analytic[i]))) err_e = float(jnp.max(jnp.abs(traj_explicit[i, :, 0] - traj_analytic[i]))) print( f"{float(mus[i]):5.1f} {imp_bounded!s:>26} {exp_bounded!s:>26} {err_i:14.4e} {err_e:14.4e}" ) # %% Plots fig, axes = plt.subplots(2, 3, figsize=(15, 8), sharex=True) for i, ax in enumerate(axes.ravel()): ax.plot(t_axis, traj_analytic[i], "k-", label="analytic") ax.plot(t_axis, traj_explicit[i, :, 0], "C0--", label="ExplicitEuler") ax.plot(t_axis, traj_implicit[i, :, 0], "C1--", label="ImplicitEuler") ax.set_title(f"mu={mus[i]:.1f}") ax.set_xlabel("t") if i == 0: ax.set_ylabel("x(t)") ax.legend() plt.tight_layout() plt.show() # %%