r"""Learns the solution of a linearized Euler equation in 1D from data. The goal is to learn a function from :math:`\mathbb{R} \times \mathbb{R} \to \mathbb{R}^2` using only data sampled from the exact solution, without any physics-based residuals. This example demonstrates function approximation from data using a data-driven projection method with a custom DataResidual. """ import timeit import jax import jax.numpy as jnp import matplotlib.pyplot as plt from scimba_jax.domains.meshless_domains.domains_1d import Segment1D from scimba_jax.nonlinear_approximation.approximation_spaces.approximation_spaces import ( ApproximationSpace, ) from scimba_jax.nonlinear_approximation.integration.monte_carlo import ( DomainSampler, TensorizedSampler, ) from scimba_jax.nonlinear_approximation.integration.monte_carlo_time import ( UniformTimeSampler, ) from scimba_jax.nonlinear_approximation.networks.mlp import MLP from scimba_jax.nonlinear_approximation.numerical_solvers.projectors import Projector from scimba_jax.physical_models.abstract_physical_model import ( AbstractPhysicalModel, ) from scimba_jax.physical_models.data_residuals import CollocDataResidual from scimba_jax.plots.plots_nd import ( plot_abstract_approx_spaces, ) N_EPOCHS = 50 N_DATA = 5000 def exact_solution(t, x): D = 0.02 coeff = 1 / (4 * jnp.pi * D) ** 0.5 p_plus_u = coeff * jnp.exp(-((x - t - 1) ** 2) / (4 * D)) p_minus_u = coeff * jnp.exp(-((x + t - 1) ** 2) / (4 * D)) p = (p_plus_u + p_minus_u) / 2 u = (p_plus_u - p_minus_u) / 2 return jnp.concatenate((p, u), axis=-1) class LinearizedEulerDataModel(AbstractPhysicalModel): """Physical model with only data residuals (no PDE residuals).""" def __init__(self, main_domain, time_domain, data): super().__init__(main_domain=main_domain, time_domain=time_domain) self.add_data_residual( "data", CollocDataResidual(size=2, model_type="t_x", data=data), ) # Define domains dx = Segment1D((-1.0, 3.0), is_main_domain=True) domain_t = (0.0, 0.5) # Generate training data key = jax.random.PRNGKey(0) key_data, key = jax.random.split(key) t_min, t_max = domain_t x_min, x_max = dx.bounds[0, 0].item(), dx.bounds[0, 1].item() # Create vectors of N_DATA points for t and x sampled uniformly key_t, key_x = jax.random.split(key_data, 2) t_data = jax.random.uniform(key_t, shape=(N_DATA, 1), minval=t_min, maxval=t_max) x_data = jax.random.uniform(key_x, shape=(N_DATA, 1), minval=x_min, maxval=x_max) # Compute exact solution at these points y_data = exact_solution(t_data, x_data) # Shape: (N_DATA, 2) print(f"Generated {N_DATA} data points with shape {y_data.shape}") # Create model model = LinearizedEulerDataModel(dx, domain_t, (t_data, x_data, y_data)) sampler = TensorizedSampler( [ UniformTimeSampler(domain_t), DomainSampler(dx), ], model_type="t_x", bc=False, ic=False, data_samplers=model.data_residuals, ) # Create approximation space nn = MLP(in_size=2, out_size=2, hidden_sizes=[16, 16], key=key) space = ApproximationSpace({"x": 1}, [(nn, "vec", 2)], model_type="t_x") # Create projector and train projector = Projector(model, space, sampler) print("\n@@@@@@@@@@@@@@@ Training from data @@@@@@@@@@@@@@@@@@@@@@") start = timeit.default_timer() key, projector = projector.project(key, space, N_EPOCHS, n_dl_colloc=1000) end = timeit.default_timer() print("best loss: ", projector.best_loss) print("time for %d epochs: " % N_EPOCHS, end - start) # Evaluate and plot print("\n@@@@@@@@@@@@@@@ Evaluation @@@@@@@@@@@@@@@@@@@@@@") plot_abstract_approx_spaces( (projector.space,), dx, time_domains=domain_t, components=([0, 1]), loss=(projector.losses,), solution=exact_solution, error=exact_solution, ) plt.show() # plt.savefig("linearized_euler_data_sol.png", dpi=150) # plt.close()