r"""Integrates the SIR epidemic ODE for a batch of initial conditions AND parameters (beta, gamma), using ``Rk4Flow`` + ``FlowsApproximationSpace`` as a plain numerical integrator -- no trainable network, no gradient, no Projector. Dynamics: dS/dt = -beta * S * I dI/dt = beta * S * I - gamma * I dR/dt = gamma * I u = [S, I, R] (merged, ``Rk4Flow``'s "u_mu" convention), mu = [beta, gamma] is passed through the approximation space's parameter axis, exactly like the physical variables in the pendulum example -- so batching over both initial conditions and (beta, gamma) is a single ``jax.vmap`` over all of ``integrate_sir``'s arguments. """ # %% 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 ( Rk4Flow, ) jax.config.update("jax_enable_x64", True) # %% Hyperparameters dt = 0.1 Tf = 100.0 n_steps = int(Tf / dt) # %% SIR dynamics -- analytic callable, no trainable state. def sir_vector_field(u, mu): """[dS/dt, dI/dt, dR/dt], u = [S, I, R] merged (Rk4Flow's "u_mu" convention), mu = [beta, gamma].""" S, I = u[0], u[1] # noqa: E741 beta, gamma = mu[0], mu[1] dS = -beta * S * I dI = beta * S * I - gamma * I dR = gamma * I return jnp.array([dS, dI, dR]) # %% Batch of initial conditions AND parameters (beta, gamma) N_sim = 6 key = jax.random.PRNGKey(0) key, k1, k2, k3 = jax.random.split(key, 4) I0s = jax.random.uniform(k1, (N_sim,), minval=0.01, maxval=0.05) S0s = 1.0 - I0s R0s = jnp.zeros(N_sim) betas = jax.random.uniform(k2, (N_sim,), minval=0.2, maxval=0.6) gammas = jax.random.uniform(k3, (N_sim,), minval=0.05, maxval=0.2) # %% Rk4Flow (classic 4th-order Runge-Kutta integrator) rk4_model = Rk4Flow( dim=3, flownet=sir_vector_field, dt=dt, time_dependent=False, params_dim=2 ) space_sir = FlowsApproximationSpace( state_dim=3, params_dim=2, model=rk4_model, model_type="u_mu", rollout=1 ) def integrate_sir(S0, I0, R0, beta, gamma): x0 = jnp.array([S0, I0, R0]) mu = jnp.array([beta, gamma]) return space_sir.rollout_trajectory(space_sir, x0, mu, n_steps) traj_sir = jax.vmap(integrate_sir)( S0s, I0s, R0s, betas, gammas ) # (N_sim, n_steps+1, 3) # %% Sanity check: S + I + R is conserved exactly by construction (every # stage of F sums to 0), regardless of the RK4 discretization -- a good # correctness check independent of the model's actual accuracy. total_population = traj_sir.sum(axis=-1) # (N_sim, n_steps + 1) print( "max |S+I+R - 1| over the batch:", float(jnp.max(jnp.abs(total_population - 1.0))), ) # %% Plots t_axis = jnp.linspace(0, n_steps * dt, n_steps + 1) fig, axes = plt.subplots(1, N_sim, figsize=(4 * N_sim, 4), sharey=True) for i, ax in enumerate(axes): ax.plot(t_axis, traj_sir[i, :, 0], label="S") ax.plot(t_axis, traj_sir[i, :, 1], label="I") ax.plot(t_axis, traj_sir[i, :, 2], label="R") ax.set_xlabel("t") ax.set_title(f"beta={betas[i]:.2f}, gamma={gammas[i]:.2f}") if i == 0: ax.set_ylabel("fraction of population") ax.legend() plt.tight_layout() plt.show() # %%