r"""Identifies the SIR contact rate ``beta`` from data. The vector field F is a ``ScimbaPytree`` holding ``beta`` as its only trainable parameter (``gamma`` stays a known, per-simulation conditioning parameter carried in ``mu`` -- see ``sir_batch_integration.py``). Wrapped in an ``ApproximationSpace`` and given to ``Rk2Flow`` as ``flownet``, exactly like an MLP flownet would be (see ``basic_discrete_ode_nets.py``): the only difference is that this "network" has a single scalar weight with an actual physical meaning. Two trainings are compared, both fed with data batched over several initial conditions (S0, I0, R0) and several gamma values but a SINGLE true beta: - rollout K=1 (single-step supervision) - rollout K=10 (dense multi-step rollout supervision, every intermediate step compared to the reference -- see pendulum_flow_rollout.py) """ # %% import timeit import jax import jax.numpy as jnp import matplotlib.pyplot as plt from scimba_jax.domains.meshless_domains.domains_nd import HypercubeND 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.nonlinear_approximation.integration.monte_carlo import ( DomainSampler, TensorizedSampler, ) from scimba_jax.nonlinear_approximation.integration.monte_carlo_parameters import ( UniformParametricSampler, ) from scimba_jax.nonlinear_approximation.numerical_solvers.projectors import Projector from scimba_jax.ode_approx.basic_discrete_ode_nets import ( Rk2Flow, Rk4Flow, ) from scimba_jax.physical_models.abstract_physical_model import ( PHYSICAL_RESIDUALS_TYPE, AbstractPhysicalModel, ) from scimba_jax.physical_models.abstract_residuals import PARAM_FUNC_TYPE, DataResidual from scimba_jax.utils.scimba_pytree import ScimbaPytree, trainable from scimba_jax.utils.typing_protocols import NDARRAY_TYPE jax.config.update("jax_enable_x64", True) # %% Hyperparameters key = jax.random.PRNGKey(0) dt = 0.1 beta_true = 0.5 beta_init = 1.0 # starting guess for beta, deliberately far from beta_true N_sim = 20 # number of (S0, I0, R0, gamma) simulations, batched together Nt_train = 200 # number of training pairs per simulation N_COLLOC = 1200 N_EPOCHS = 5 ENG_REGULARIZATION = 1e-6 # %% SIR vector field -- a ScimbaPytree with a single trainable parameter. class SIRVectorField(ScimbaPytree): """F(u, mu) for the SIR model, u = [S, I, R] merged, mu = [gamma]. ``beta`` is declared ``trainable(True)``, which is what makes it a parameter: the default is FROZEN, so a field is optimised only because someone said it should be. Being an array is not enough and is not supposed to be -- the person writing the model is the one who knows what is being identified. (This model would in fact learn without the marker, since an ``ApproximationSpace`` declares its own models trainable and the role is inherited; the declaration says it in the one place a reader looks.) Called as ``self(inputs)`` with ``inputs = [S, I, R, gamma]`` concatenated (``ApproximationSpace``'s default pre-processing), matching the raw-array convention any "vec" model must follow. """ beta: NDARRAY_TYPE = trainable(True) def __init__(self, beta_init: float): self.beta = jnp.array([beta_init]) def __call__(self, inputs: NDARRAY_TYPE) -> NDARRAY_TYPE: S, I, gamma = inputs[0], inputs[1], inputs[3] # noqa: E741 beta = self.beta[0] dS = -beta * S * I dI = beta * S * I - gamma * I dR = gamma * I return jnp.array([dS, dI, dR]) def ndof(self) -> int: return 1 # %% Reference SIR dynamics (true beta) -- plain analytic callable, no # trainable state, plugged into the same Rk4Flow used elsewhere # (sir_batch_integration.py) instead of a hand-rolled RK4 stepper. def sir_true_vector_field(u, mu): """[dS/dt, dI/dt, dR/dt], u = [S, I, R] merged, mu = [gamma] -- beta is fixed to the (known) true value here, closed over as a constant.""" S, I = u[0], u[1] # noqa: E741 gamma = mu[0] dS = -beta_true * S * I dI = beta_true * S * I - gamma * I dR = gamma * I return jnp.array([dS, dI, dR]) reference_model = Rk4Flow( dim=3, flownet=sir_true_vector_field, dt=dt, time_dependent=False, params_dim=1 ) reference_space = FlowsApproximationSpace( state_dim=3, params_dim=1, model=reference_model, model_type="u_mu", rollout=1 ) # %% Batch of initial conditions AND gamma (single shared true beta) 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) gammas = jax.random.uniform(k2, (N_sim,), minval=0.05, maxval=0.2) def make_training_data(n_rollout): """Generate pairs (x_t, [x_{t+dt}, ..., x_{t+n_rollout*dt}]) for all simulations: one dense target per model-call checkpoint, matching what FlowsApproximationSpace.create_variables() returns (the flat rollout trajectory) -- see pendulum_flow_rollout.py. """ def simulate_one(S0, I0, R0, gamma): x0 = jnp.array([S0, I0, R0]) mu = jnp.array([gamma]) # x_all[i] = state at t = i * dt, starting with x0 (see rollout_trajectory). x_all = reference_space.rollout_trajectory( reference_space, x0, mu, Nt_train + n_rollout - 1 ) x_start = x_all[:Nt_train] offsets = jnp.arange(1, n_rollout + 1) targets = jnp.stack( [x_all[off : off + Nt_train] for off in offsets], axis=1 ) # (Nt_train, n_rollout, 3) return x_start, targets x_start, targets = jax.vmap(simulate_one)(S0s, I0s, R0s, gammas) # x_start: (N_sim, Nt_train, 3); targets: (N_sim, Nt_train, n_rollout, 3) x = x_start.reshape(-1, 3) mu = jnp.repeat(gammas, Nt_train)[:, None] y = targets.reshape(-1, n_rollout * 3) # flat, step-major (see self_compose_traj) return x, mu, y # %% Infrastructure: domain, sampler, model, residual class ModelEmpty(AbstractPhysicalModel): def __init__(self, main_domain): super().__init__(main_domain=main_domain) self.physical_residuals: PHYSICAL_RESIDUALS_TYPE = {} class SIRDataResidual(DataResidual): def __init__( self, size: int, model_type: str = "u_mu", data: tuple[jnp.ndarray, ...] = tuple(), ): super().__init__(size=size, model_type=model_type, data=data) def construct_residual(self, *vars: PARAM_FUNC_TYPE) -> PARAM_FUNC_TYPE: return vars[0] # forward flow variable (flat rollout trajectory) domain_x = HypercubeND([(0.0, 1.0), (0.0, 1.0), (0.0, 1.0)], is_main_domain=True) domain_mu = [(0.05, 0.2)] def make_sampler(model): return TensorizedSampler( [DomainSampler(domain_x), UniformParametricSampler(domain_mu)], bc=False, data_samplers=model.data_residuals, ) def make_pde_model(n_rollout, data): model = ModelEmpty(main_domain=domain_x) model.add_data_residual( "data", SIRDataResidual(size=n_rollout * 3, model_type="u_mu", data=data) ) return model def train_sir(key, n_rollout, beta_init=beta_init): x_data, mu_data, y_data = make_training_data(n_rollout) pde_model = make_pde_model(n_rollout, (x_data, mu_data, y_data)) sampler = make_sampler(pde_model) sir_field = SIRVectorField(beta_init=beta_init) sir_field_space = ApproximationSpace( dims={"u": 3, "mu": 1}, list_models=[(sir_field, "vec", 3)], model_type="u_mu" ) rk2_model = Rk2Flow( dim=3, flownet=sir_field_space, dt=dt, time_dependent=False, params_dim=1 ) space = FlowsApproximationSpace( state_dim=3, params_dim=1, model=rk2_model, model_type="u_mu", rollout=n_rollout ) pinn = Projector( pde_model, space, sampler, matrix_regularization=ENG_REGULARIZATION ) key, pinn = pinn.project(key, space, N_EPOCHS, N_COLLOC) beta_learned = float(pinn.space.models[0].F.models[0].beta[0]) return key, pinn, beta_learned # %% Train with K=1 (single-step) and K=10 (multi-step rollout) print(f"true beta = {beta_true}") start = timeit.default_timer() key, pinn_k1, beta_k1 = train_sir(key, n_rollout=1) print( f"K=1: learned beta = {beta_k1:.4f} " f"(final loss {pinn_k1.best_loss['total']:.4e}, {timeit.default_timer() - start:.1f}s)" ) start = timeit.default_timer() key, pinn_k10, beta_k10 = train_sir(key, n_rollout=10) print( f"K=10: learned beta = {beta_k10:.4f} " f"(final loss {pinn_k10.best_loss['total']:.4e}, {timeit.default_timer() - start:.1f}s)" ) # %% Plots fig, axes = plt.subplots(1, 2, figsize=(12, 5)) ax = axes[0] ax.semilogy(pinn_k1.losses.losses_history["total"], label="K=1") ax.semilogy(pinn_k10.losses.losses_history["total"], label="K=10") ax.set_xlabel("epoch") ax.set_ylabel("loss") ax.set_title("Training loss") ax.legend() ax = axes[1] ax.bar( ["true", "init", "K=1", "K=10"], [beta_true, beta_init, beta_k1, beta_k10], color=["black", "gray", "C0", "C1"], ) ax.set_ylabel("beta") ax.set_title("Identified beta") plt.tight_layout() plt.show() # %%