r"""Learns a metric gradient flow ``dx/dt = -G grad E(x)`` from trajectory data. The reference dynamics is a *true* metric gradient flow ``F = -G grad E_true`` with a known quadratic energy ``E_true(x) = 1/2 (x0^2 + 2 x1^2)`` and a fixed symmetric positive-definite metric ``G`` (so ``F_true = -G [x0, 2 x1]``), integrated with ``Rk4Flow`` to produce reference trajectories. The learned model wraps a trainable MLP energy ``E_theta`` in an ``ApproximationSpace`` and feeds it to :class:`GradientFlowVectorFieldSpace` -- which exposes ``F = -G grad E_theta`` with the same (fixed, known) metric ``G`` -- plugged into an explicit ``Rk2Flow`` exactly like any other flownet. Only the energy MLP is trained (gradients flow through ``GradientFlowVectorFieldSpace`` into the MLP weights), by fitting the K-step rollout to the reference, via the same ``Projector`` / ``DataResidual`` machinery as ``sir_beta_identification.py``. The learned vector field only ever sees ``-G grad E_theta``, so it can only represent (and is regularized towards) genuine metric gradient dynamics -- the energy is recovered up to an additive constant. """ 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.networks.mlp import MLP 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.ode_approx.vector_fields import ( GradientFlowVectorFieldSpace, ) 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 jax.config.update("jax_enable_x64", True) # %% Hyperparameters key = jax.random.PRNGKey(0) dt = 0.05 N_sim = 16 # number of initial conditions, batched together Nt_train = 60 # number of training pairs per trajectory N_ROLLOUT = 5 # multi-step rollout supervision (K) N_COLLOC = 1500 N_EPOCHS = 60 ENG_REGULARIZATION = 1e-6 # %% Reference: the TRUE gradient flow F = -grad E_true (analytic callable). # A constant symmetric positive-definite metric G: F = -G grad E. G_METRIC = jnp.array([[1.5, 0.5], [0.5, 1.0]]) def energy_true(u): """Quadratic energy E_true(x) = 1/2 (x0^2 + 2 x1^2).""" return 0.5 * (u[0] ** 2 + 2.0 * u[1] ** 2) def gradient_flow_true(u, mu): """F_true = -G grad E_true = -G [x0, 2 x1] (mu unused, params_dim = 0).""" return -G_METRIC @ jnp.array([u[0], 2.0 * u[1]]) def metric_fn(u, mu): """The (fixed, known) metric G as a flat (dim*dim,) vector.""" return jnp.reshape(G_METRIC, (4,)) reference_model = Rk4Flow( dim=2, flownet=gradient_flow_true, dt=dt, time_dependent=False, params_dim=0 ) reference_space = FlowsApproximationSpace( state_dim=2, params_dim=0, model=reference_model, model_type="u_mu", rollout=1 ) # %% Batch of initial conditions key, k0 = jax.random.split(key) X0 = jax.random.uniform(k0, (N_sim, 2), minval=-1.5, maxval=1.5) def make_training_data(n_rollout): """(x_t, [x_{t+dt}, ..., x_{t+K dt}]) pairs from the reference gradient flow.""" def simulate_one(x0): x_all = reference_space.rollout_trajectory( reference_space, x0, jnp.array([]), 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, 2) return x_start, targets x_start, targets = jax.vmap(simulate_one)(X0) x = x_start.reshape(-1, 2) mu = jnp.zeros((x.shape[0], 0)) y = targets.reshape(-1, n_rollout * 2) # flat, step-major return x, mu, y # %% Infrastructure: empty PDE model + a data residual on the rollout trajectory class ModelEmpty(AbstractPhysicalModel): def __init__(self, main_domain): super().__init__(main_domain=main_domain) self.physical_residuals: PHYSICAL_RESIDUALS_TYPE = {} class RolloutDataResidual(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] # the flat K-step rollout trajectory domain_x = HypercubeND([(-2.0, 2.0), (-2.0, 2.0)], is_main_domain=True) def make_sampler(model): return TensorizedSampler( [DomainSampler(domain_x), UniformParametricSampler([])], 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", RolloutDataResidual(size=n_rollout * 2, model_type="u_mu", data=data) ) return model # %% Build the dynamic gradient-flow space and train it def build_space(key, n_rollout): """A gradient-flow flownet with a trainable MLP energy, in an Rk2 flow.""" key, sub = jax.random.split(key) energy_net = MLP(in_size=2, out_size=1, hidden_sizes=[32, 32], key=sub) energy_space = ApproximationSpace( dims={"u": 2, "mu": 0}, list_models=[(energy_net, "scalar", None)], model_type="u_mu", ) vf = GradientFlowVectorFieldSpace( dim=2, energy=energy_space, metric=metric_fn, params_dim=0 ) rk2_model = Rk2Flow(dim=2, flownet=vf, dt=dt, time_dependent=False, params_dim=0) space = FlowsApproximationSpace( state_dim=2, params_dim=0, model=rk2_model, model_type="u_mu", rollout=n_rollout ) return key, space def rollout_error(space, n_steps=Nt_train): """Max abs error of a learned trajectory vs the reference, over the batch.""" ref = jax.vmap( lambda x0: reference_space.rollout_trajectory( reference_space, x0, jnp.array([]), n_steps ) )(X0) got = jax.vmap( lambda x0: space.rollout_trajectory(space, x0, jnp.array([]), n_steps) )(X0) return float(jnp.max(jnp.abs(got - ref))) 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) key, space = build_space(key, N_ROLLOUT) # a rollout-1 view of the same trained model, for the trajectory error metric def rollout1_view(space): return FlowsApproximationSpace( state_dim=2, params_dim=0, model=space.models[0], model_type="u_mu", rollout=1, ) err_before = rollout_error(rollout1_view(space)) print(f"trajectory max error BEFORE training = {err_before:.4e}") pinn = Projector(pde_model, space, sampler, matrix_regularization=ENG_REGULARIZATION) key, pinn = pinn.project(key, space, N_EPOCHS, N_COLLOC) err_after = rollout_error(rollout1_view(pinn.space)) print(f"trajectory max error AFTER training = {err_after:.4e}") print(f"training best loss = {pinn.best_loss['total']:.3e}") # %% Plot: learned vs reference trajectory from a held-out initial condition key, kt = jax.random.split(key) x0_test = jax.random.uniform(kt, (2,), minval=-1.5, maxval=1.5) n_plot = 80 ref_traj = reference_space.rollout_trajectory( reference_space, x0_test, jnp.array([]), n_plot ) learned_traj = rollout1_view(pinn.space).rollout_trajectory( rollout1_view(pinn.space), x0_test, jnp.array([]), n_plot ) ref_traj = jax.device_get(ref_traj) learned_traj = jax.device_get(learned_traj) fig, ax = plt.subplots(figsize=(6, 6)) ax.plot(ref_traj[:, 0], ref_traj[:, 1], "k-", lw=2, label="reference (-G grad E_true)") ax.plot( learned_traj[:, 0], learned_traj[:, 1], "r--", lw=2, label="learned (-G grad E_theta)", ) ax.plot(x0_test[0], x0_test[1], "go", label="x0") ax.set_xlabel("x0") ax.set_ylabel("x1") ax.set_aspect("equal") ax.legend() ax.set_title("Gradient-flow learning: learned energy vs true energy") plt.tight_layout() plt.savefig("gradient_flow_learning.png", dpi=130, bbox_inches="tight") print("Saved: gradient_flow_learning.png")