r"""PINN on the flow of a pendulum ODE: a space-time network, not a solver. Dynamics given by the Hamiltonian: H(q, p, mu) = p²/2 + mu*q²/2 + mu*0.012*q³/3 The flow phi^0(t, u0) of the ODE u' = F(u, mu), F = [dH/dp, -dH/dq], satisfies (see e.g. Def. "Flot d'une EDO"): d/dt phi^0(t, u0, mu) = F(phi^0(t, u0, mu), mu), phi^0(0, u0, mu) = u0 We DON'T use a numerical integrator (Rk4Flow, ...) to represent the trained model: a single MLP u_theta(t, u0, mu) -> (q, p) directly represents the flow map itself, for any t and any initial condition u0=(q0, p0). Rk4Flow is only used, separately, to generate synthetic reference/training trajectories (the "ground truth" data) -- never inside the trained model. Physics residual (made explicit, via autodiff of u_theta w.r.t. t): R(t, u0, mu) = d(u_theta)/dt(t, u0, mu) - F(u_theta(t, u0, mu), mu) = [dq_theta/dt - p_theta , dp_theta/dt + mu*q_theta + 0.012*mu*q_theta^2] sampled on broad (t, u0, mu) collocation points. Combined with a data loss fitting u_theta(t_i, u0_i, mu_i) to sampled reference-trajectory points, in a single Projector. Two cases are compared: data loss alone vs. physics+data. """ # %% import timeit import jax import jax.numpy as jnp import matplotlib.pyplot as plt from scimba_jax.nonlinear_approximation.approximation_spaces.approximation_spaces import ( AbstractApproxSpace, ApproximationSpace, ) from scimba_jax.nonlinear_approximation.approximation_spaces.flow_approximation_spaces import ( FlowsApproximationSpace, ) from scimba_jax.nonlinear_approximation.integration.data_sampler import DataSampler from scimba_jax.nonlinear_approximation.integration.monte_carlo import ( TensorizedSampler, ) from scimba_jax.nonlinear_approximation.integration.monte_carlo_parameters import ( UniformParametricSampler, ) from scimba_jax.nonlinear_approximation.integration.monte_carlo_time import ( UniformTimeSampler, ) from scimba_jax.nonlinear_approximation.model_class.funcparam_matrix import ( ParamVecFunction, arg, ) from scimba_jax.nonlinear_approximation.networks.mlp import MLP from scimba_jax.nonlinear_approximation.numerical_solvers.projectors import Projector from scimba_jax.ode_approx.basic_discrete_ode_nets import ( 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, InteriorResidual, ) jax.config.update("jax_enable_x64", True) # %% Hyperparameters key = jax.random.PRNGKey(0) dt = 0.03 # Time step used only to generate reference/training trajectories N_simu = 50 # Number of reference trajectories Nt_train = 200 # Time steps per trajectory -> T_train = Nt_train * dt = 1.5 N_COLLOC = 2000 # Collocation points for physics loss N_EPOCHS = 200 # Reduced for demo; increase to 200+ for production # %% Pendulum dynamics def hamiltonian(q, p, mu): """H(q, p) = p²/2 + mu*q²/2 + mu*0.012*q³/3""" return p**2 / 2.0 + mu * q**2 / 2.0 + mu * 0.012 * q**3 / 3.0 def canonical_vf(u, mu): q, p = u[0], u[1] dh_dq = mu[0] * q + mu[0] * 0.012 * q**2 dh_dp = p return jnp.array([dh_dp, -dh_dq]) # [dq/dt, dp/dt] # %% Generate reference trajectories with scimba's Rk4Flow -- ground truth # data only, NOT part of the trained model (which is a pure space-time net). print("\n=== Generate reference trajectories via scimba Rk4Flow (ground truth) ===") analytical_flow = Rk4Flow( dim=2, flownet=canonical_vf, dt=dt, time_dependent=False, params_dim=1 ) analytical_space = FlowsApproximationSpace( state_dim=2, params_dim=1, model=analytical_flow, model_type="u_mu", rollout=1 ) key, k1, k2, k3 = jax.random.split(key, 4) q0 = 1.0 + jax.random.uniform(k1, (N_simu,)) * 2.0 # q0 in [1, 3] p0 = 0.1 * jax.random.uniform(k2, (N_simu,)) * 2.0 # p0 in [0, 0.2] mu_simu = 0.8 + jax.random.uniform(k3, (N_simu,)) # mu in [0.8, 1.8] def simulate_one(q0_i, p0_i, mu_i): u0 = jnp.array([q0_i, p0_i]) mu0 = jnp.array([mu_i]) return analytical_space.rollout_trajectory(analytical_space, u0, mu0, Nt_train) trajs = jax.vmap(simulate_one)(q0, p0, mu_simu) # (N_simu, Nt_train + 1, 2) # Data for the space-time network: (t, u0, mu) -> (q, p) observed along each # reference trajectory (sim-major order, matching trajs.reshape below). t_grid = jnp.arange(Nt_train + 1) * dt # (Nt_train + 1,) t_data = jnp.tile(t_grid, N_simu)[:, None] # (N, 1) u0_data = jnp.repeat(jnp.stack([q0, p0], axis=-1), Nt_train + 1, axis=0) # (N, 2) mu_data = jnp.repeat(mu_simu, Nt_train + 1)[:, None] # (N, 1) y_data = trajs.reshape(-1, 2) # (N, 2), observed (q, p) print(f"Training data: {t_data.shape[0]} (t, u0, mu) -> (q, p) samples") # %% Reference trajectory for validation/plots (analytical Rk4 rollout) domain_t_end = 10.0 # total trained time horizon [0, T] Nt_ref = int(domain_t_end / dt) # cover the whole trained horizon [0, T] q0ref, p0ref, mu_ref = 1.4, 0.15, 0.7 u0_ref = jnp.array([q0ref, p0ref]) mu0_ref = jnp.array([mu_ref]) traj_ref = analytical_space.rollout_trajectory( analytical_space, u0_ref, mu0_ref, Nt_ref ) q_ref, p_ref = traj_ref[:, 0], traj_ref[:, 1] h_ref = jax.vmap(lambda q, p: hamiltonian(q, p, mu_ref))(q_ref, p_ref) # %% Infrastructure: sampler, data residual (shared by both cases) domain_t = (0.0, domain_t_end) domain_u0 = [(1.0, 3.0), (0.0, 0.2)] domain_mu = [(0.8, 1.8)] data_sampler = DataSampler((t_data, u0_data, mu_data, y_data)) sampler = TensorizedSampler( [ UniformTimeSampler(domain_t), UniformParametricSampler(domain_u0), UniformParametricSampler(domain_mu), ], model_type="t_u0_mu", bc=False, ic=False, data_samplers={"data": data_sampler}, domain_vars=[], # no bounded/boundary domain: u0 is just a free parameter ) class ModelEmpty(AbstractPhysicalModel): def __init__(self, time_domain): super().__init__(main_domain=None, time_domain=time_domain) self.physical_residuals: PHYSICAL_RESIDUALS_TYPE = {} class TrajDataResidual(DataResidual): """Data residual: u_theta(t, u0, mu) -> (q, p), matched to observed data.""" def __init__(self, model_type: str = "t_u0_mu"): super().__init__(size=2, model_type=model_type) def construct_residual(self, *vars: PARAM_FUNC_TYPE) -> PARAM_FUNC_TYPE: return vars[0] # predicted (q_theta, p_theta) pde_model_base = ModelEmpty(time_domain=domain_t) pde_model_base.add_data_residual("data", TrajDataResidual()) # %% Physics residual: Hamilton's canonical equations on the flow itself. # # d(u_theta)/dt is obtained by autodiff of the network w.r.t. its OWN "t" # input (``ParamFunction.d_t()``, a jacrev), then compared to F evaluated at # the network's OWN predicted state (q_theta, p_theta) -- this literally is # the flow equation phi'(t) = F(phi(t), mu), not a comparison against an # externally supplied target. class HamiltonianODEResidual(InteriorResidual): def __init__(self, time_domain): super().__init__(size=2, model_type="t_u0_mu", time_domain=time_domain) def construct_residual(self, *vars: PARAM_FUNC_TYPE) -> PARAM_FUNC_TYPE: u = vars[0] # (q_theta, p_theta) as a function of (t, u0, mu) dudt = u.d_t() # d/dt (q_theta, p_theta), via autodiff on t q, p = u.components() mu = arg(u, "mu").component(0) dh_dq = mu * q + 0.012 * mu * q * q # ∂H/∂q at (q_theta, p_theta) dh_dp = p # ∂H/∂p at (q_theta, p_theta) F_true = ParamVecFunction.cat([dh_dp, -dh_dq]) return dudt - F_true # residual, compared to f_rhs = 0 class ModelPhysicsAndData(AbstractPhysicalModel): def __init__(self, time_domain): super().__init__(main_domain=None, time_domain=time_domain) self.physical_residuals: PHYSICAL_RESIDUALS_TYPE = { "interior": HamiltonianODEResidual(time_domain=time_domain) } def renew_time_domain(self, new_time_domain): """Restrict the model + its residuals to a new time sub-interval.""" self.time_domain = new_time_domain for label in self.physical_residuals: self.physical_residuals[label].time_domain = new_time_domain pde_model_physics = ModelPhysicsAndData(time_domain=domain_t) pde_model_physics.add_data_residual("data", TrajDataResidual()) # %% ComposedFlowSpace: chain N continuous flows in time, one trained at a time. # # Same pattern as scimba's ``DensityFlowApproximationSpace`` (multiflow): # ``idx_current_flow`` selects the single sub-flow currently being trained -- # only ``models[idx_current_flow]`` is a pytree child (so only its ~2000 dof # get gradients); all other sub-flows go to ``frozen_models`` in aux_data # (constants, no gradient). We train phi_0 on [T_0, T_1], then freeze it and # train phi_1 on [T_1, T_2], etc. The composed flow phi(t, u0, mu) for # t in [T_k, T_{k+1}] applies the frozen phi_0..phi_{k-1} over their full # sub-intervals to carry u0 up to T_k, then the current phi_k on [T_k, t]. class ComposedFlowSpace(AbstractApproxSpace): """Sequentially-trained composition of continuous flows (phi_k on [T_k, T_{k+1}]).""" idx_current_flow: int = 0 time_intervals: tuple = () def __init__( self, sub_spaces: list[AbstractApproxSpace], time_intervals: list[tuple[float, float]], ): super().__init__( dims={"t": 1, "u0": 2, "mu": 1}, model_type="t_u0_mu", models=list(sub_spaces), types_models=["vec"] * len(sub_spaces), size_models=[2] * len(sub_spaces), ) self.idx_current_flow = 0 self.time_intervals = tuple(time_intervals) def compute_ndof(self) -> int: # Only the flow currently being trained carries trainable dof. return self.models[self.idx_current_flow].compute_ndof() def create_variables(self) -> tuple[ParamVecFunction, ...]: idx = self.idx_current_flow intervals = self.time_intervals def fn(space: "ComposedFlowSpace", t, u0, mu): u = u0 # Frozen flows 0..idx-1: carry u0 up to T_idx over full sub-intervals. for j in range(idx): phi_j = space.models[j].create_variables()[0] dt_j = jnp.array([intervals[j][1] - intervals[j][0]]) u = phi_j(space.models[j], dt_j, u, mu) # Current flow idx: local time t - T_idx, from the carried-in state. phi_idx = space.models[idx].create_variables()[0] t_local = t - jnp.array([intervals[idx][0]]) return phi_idx(space.models[idx], t_local, u, mu) return (ParamVecFunction(2, self.dims, fn, f_type="t_u0_mu"),) def tree_flatten(self): _, parent_aux = super().tree_flatten() idx = self.idx_current_flow aux = { "idx_current_flow": idx, "time_intervals": self.time_intervals, "frozen_models": list(self.models[:idx]) + list(self.models[idx + 1 :]), "parent_aux_data": parent_aux, } return [self.models[idx]], aux @classmethod def tree_unflatten(cls, aux, children): obj = super(ComposedFlowSpace, cls).tree_unflatten( aux["parent_aux_data"], children ) idx = aux["idx_current_flow"] obj.idx_current_flow = idx obj.time_intervals = aux["time_intervals"] frozen = list(aux["frozen_models"]) obj.models = frozen[:idx] + list(children) + frozen[idx:] return obj # %% Case 1: space-time MLP, data loss only print("\n#1: space-time MLP u_theta(t, u0, mu) -> (q, p)") print(" Loss: data only") key, k1 = jax.random.split(key) nn1 = MLP(in_size=1 + 2 + 1, out_size=2, hidden_sizes=[32, 32, 32], key=k1) space1 = ApproximationSpace( dims={"t": 1, "u0": 2, "mu": 1}, list_models=[(nn1, "vec", 2)], model_type="t_u0_mu" ) print(f" ndof: {space1.compute_ndof()}") pinn1 = Projector(pde_model_base, space1, sampler) start = timeit.default_timer() key, pinn1 = pinn1.project(key, space1, N_EPOCHS, N_COLLOC) print( f" final loss: {pinn1.best_loss['total']:.4e} ({timeit.default_timer() - start:.1f}s)" ) space1 = pinn1.space # %% Case 2: same architecture, physics (Hamilton's equations) + data loss print("\n#2: space-time MLP, same architecture as #1") print(" Loss: physics (Hamilton's equations, autodiff on t) + data") key, k2 = jax.random.split(key) nn2 = MLP(in_size=1 + 2 + 1, out_size=2, hidden_sizes=[32, 32, 32], key=k2) space2 = ApproximationSpace( dims={"t": 1, "u0": 2, "mu": 1}, list_models=[(nn2, "vec", 2)], model_type="t_u0_mu" ) print(f" ndof: {space2.compute_ndof()}") pinn2 = Projector( pde_model_physics, space2, sampler, weights={"interior": [1.0, 1.0], "data": [10.0, 10.0]}, ) start = timeit.default_timer() key, pinn2 = pinn2.project(key, space2, N_EPOCHS, N_COLLOC) print( f" final loss: {pinn2.best_loss['total']:.4e} ({timeit.default_timer() - start:.1f}s)" ) print(f" physics (interior): {jnp.sum(pinn2.best_loss['interior']):.4e}") print(f" data: {jnp.sum(pinn2.best_loss['data']):.4e}") space2 = pinn2.space # %% Case 3: same physics+data loss as #2, but 2 flows trained SEQUENTIALLY. # phi_0 on [0, T1] first; then freeze it and train phi_1 on [T1, T2]. Each # stage trains only ~2000 dof (one sub-flow), never both at once. print("\n#3: composed flow, 2 sub-flows trained sequentially (phi_0 then phi_1)") print(" Loss: physics (Hamilton's equations, autodiff on t) + data") T_split = domain_t[1] / 2.0 time_intervals = [(0.0, T_split), (T_split, domain_t[1])] key, k3a, k3b = jax.random.split(key, 3) nn3a = MLP(in_size=1 + 2 + 1, out_size=2, hidden_sizes=[32, 32, 32], key=k3a) nn3b = MLP(in_size=1 + 2 + 1, out_size=2, hidden_sizes=[32, 32, 32], key=k3b) space3a = ApproximationSpace( dims={"t": 1, "u0": 2, "mu": 1}, list_models=[(nn3a, "vec", 2)], model_type="t_u0_mu", ) space3b = ApproximationSpace( dims={"t": 1, "u0": 2, "mu": 1}, list_models=[(nn3b, "vec", 2)], model_type="t_u0_mu", ) space3 = ComposedFlowSpace([space3a, space3b], time_intervals) pinn3 = None for i, interval in enumerate(time_intervals): space3.idx_current_flow = i pde_model_physics.renew_time_domain(interval) # Sampler restricted to this sub-interval: collocation time in `interval`, # data filtered to the same window. mask = (t_data[:, 0] >= interval[0]) & (t_data[:, 0] <= interval[1]) data_sampler_i = DataSampler( (t_data[mask], u0_data[mask], mu_data[mask], y_data[mask]) ) sampler_i = TensorizedSampler( [ UniformTimeSampler(interval), UniformParametricSampler(domain_u0), UniformParametricSampler(domain_mu), ], model_type="t_u0_mu", bc=False, ic=False, data_samplers={"data": data_sampler_i}, domain_vars=[], ) print( f" flow {i + 1}/{len(time_intervals)} on {interval}: " f"ndof={space3.compute_ndof()} (only this sub-flow trains)" ) pinn3 = Projector( pde_model_physics, space3, sampler_i, weights={"interior": [1.0, 1.0], "data": [10.0, 10.0]}, ) start = timeit.default_timer() key, pinn3 = pinn3.project(key, space3, N_EPOCHS, N_COLLOC) print( f" final loss: {pinn3.best_loss['total']:.4e} " f"physics={jnp.sum(pinn3.best_loss['interior']):.4e} " f"data={jnp.sum(pinn3.best_loss['data']):.4e} " f"({timeit.default_timer() - start:.1f}s)" ) space3 = pinn3.space space3.idx_current_flow = len(time_intervals) - 1 # full composed flow for eval # %% Evaluate: the trained space-time net IS the trajectory, no integrator. def predict_trajectory(space, t_array, u0, mu): u_pf = space.create_variables()[0] return jax.vmap(lambda t: u_pf(space, jnp.array([t]), u0, mu))(t_array) def predict_composed(space, t_array, u0, mu): """Evaluate the composed flow: for each t, use the sub-flow whose interval contains t (idx_current_flow = k selects phi_k on [T_k, T_{k+1}], with the frozen phi_0..phi_{k-1} carrying u0 up to T_k inside create_variables).""" out = jnp.zeros((t_array.shape[0], 2)) for k, (t0, t1) in enumerate(space.time_intervals): space.idx_current_flow = k u_pf = space.create_variables()[0] seg = jax.vmap(lambda t: u_pf(space, jnp.array([t]), u0, mu))(t_array) mask = (t_array >= t0) & (t_array <= t1) out = jnp.where(mask[:, None], seg, out) return out t_axis = jnp.linspace(0, Nt_ref * dt, Nt_ref + 1) traj1 = predict_trajectory(space1, t_axis, u0_ref, mu0_ref) traj2 = predict_trajectory(space2, t_axis, u0_ref, mu0_ref) traj3 = predict_composed(space3, t_axis, u0_ref, mu0_ref) q_pred1, p_pred1 = traj1[:, 0], traj1[:, 1] q_pred2, p_pred2 = traj2[:, 0], traj2[:, 1] q_pred3, p_pred3 = traj3[:, 0], traj3[:, 1] h_pred1 = jax.vmap(lambda q, p: hamiltonian(q, p, mu_ref))(q_pred1, p_pred1) h_pred2 = jax.vmap(lambda q, p: hamiltonian(q, p, mu_ref))(q_pred2, p_pred2) h_pred3 = jax.vmap(lambda q, p: hamiltonian(q, p, mu_ref))(q_pred3, p_pred3) # %% Plots fig, axes = plt.subplots(2, 2, figsize=(14, 10)) # Phase portrait ax = axes[0, 0] ax.plot( q_ref, p_ref, "k-", linewidth=2.5, label="Reference (RK4)", marker="x", markevery=Nt_ref // 20, ) ax.plot(q_pred1, p_pred1, "--", linewidth=1.5, label="Space-time NN (data only)") ax.plot(q_pred2, p_pred2, "--", linewidth=1.5, label="Space-time NN (physics + data)") ax.plot(q_pred3, p_pred3, "--", linewidth=1.5, label="Composed flow (physics + data)") ax.set_xlabel("q", fontsize=11) ax.set_ylabel("p", fontsize=11) ax.set_title("Phase portrait", fontsize=12) ax.legend() ax.grid(alpha=0.3) # Training loss ax = axes[0, 1] lh1 = pinn1.losses.losses_history lh2 = pinn2.losses.losses_history lh3 = pinn3.losses.losses_history ax.semilogy(lh1["total"], label="data only", linewidth=1.5) ax.semilogy(lh2["total"], label="physics + data, total", linewidth=1.5) ax.semilogy(lh2["interior"], "--", label=" ↳ physics component", linewidth=1.2) ax.semilogy(lh2["data"], "--", label=" ↳ data component", linewidth=1.2) ax.semilogy(lh3["total"], label="composed flow, total", linewidth=1.5) ax.set_xlabel("epoch", fontsize=11) ax.set_ylabel("loss", fontsize=11) ax.set_title("Training loss convergence", fontsize=12) ax.legend() ax.grid(alpha=0.3) # Time evolution of q ax = axes[1, 0] ax.plot(t_axis, q_ref, "k-", linewidth=2.5, label="Reference") ax.plot(t_axis, q_pred1, "--", linewidth=1.5, label="data only", alpha=0.7) ax.plot(t_axis, q_pred2, "--", linewidth=1.5, label="physics + data", alpha=0.7) ax.plot(t_axis, q_pred3, "--", linewidth=1.5, label="composed flow", alpha=0.7) ax.axvline(Nt_train * dt, color="gray", linestyle=":", label="train/extrapolate") ax.axvline(T_split, color="red", linestyle=":", label="phi1/phi2 split") ax.set_xlabel("t", fontsize=11) ax.set_ylabel("q(t)", fontsize=11) ax.set_title("Position vs time", fontsize=12) ax.legend() ax.grid(alpha=0.3) # Hamiltonian (energy) preservation ax = axes[1, 1] ax.plot(t_axis, h_ref, "k-", linewidth=2.5, label="Reference (true H)") ax.plot(t_axis, h_pred1, "--", linewidth=1.5, label="data only", alpha=0.7) ax.plot(t_axis, h_pred2, "--", linewidth=1.5, label="physics + data", alpha=0.7) ax.plot(t_axis, h_pred3, "--", linewidth=1.5, label="composed flow", alpha=0.7) ax.axvline(Nt_train * dt, color="gray", linestyle=":") ax.axvline(T_split, color="red", linestyle=":") ax.set_xlabel("t", fontsize=11) ax.set_ylabel("H(q, p)", fontsize=11) ax.set_title("Energy conservation", fontsize=12) ax.legend() ax.grid(alpha=0.3) plt.tight_layout() plt.show() # %% Diagnostics: Phase-space error and energy drift def phase_error(q_ref, p_ref, q_pred, p_pred): """L2 error in phase space.""" return jnp.sqrt(jnp.mean((q_ref - q_pred) ** 2 + (p_ref - p_pred) ** 2)) def energy_drift(h_ref, h_pred): """Max absolute energy drift.""" return jnp.max(jnp.abs(h_pred - h_ref)) print("\n=== Diagnostics ===") err1 = phase_error(q_ref, p_ref, q_pred1, p_pred1) err2 = phase_error(q_ref, p_ref, q_pred2, p_pred2) err3 = phase_error(q_ref, p_ref, q_pred3, p_pred3) drift1 = energy_drift(h_ref, h_pred1) drift2 = energy_drift(h_ref, h_pred2) drift3 = energy_drift(h_ref, h_pred3) print(f"data only: phase error={err1:.4e} max energy drift={drift1:.4e}") print(f"physics + data: phase error={err2:.4e} max energy drift={drift2:.4e}") print(f"composed flow (phys): phase error={err3:.4e} max energy drift={drift3:.4e}") # %%