r"""Solves the wave equation in 2D using a PINN. .. math:: \partial_{tt} u - \Delta u & = 0 in \Omega \times (0, T) \\ u & = u_0 on \Omega \times {0} \\ \partial_t u & = 0 on \Omega \times {0} where :math:`u: \partial \Omega \times (0, T) \to \mathbb{R}` is the unknown function, :math:`\Omega \subset \mathbb{R}^2` is the spatial domain and :math:`(0, T) \subset \mathbb{R}` is the time domain. Periodic boundary conditions are prescribed, and the initial condition is a Gaussian centered in the middle of the domain. The equation is solved on a square domain; strong boundary conditions and weak initial conditions are used. """ 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 ( 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.temporal_pde.wave_equations import WaveND from scimba_jax.plots.plots_nd import ( plot_abstract_approx_space, ) DIMENSION = 2 N_COLLOC = 1500 * (DIMENSION + 1) N_EPOCHS = 50 * 5 + 250 + 50 * (DIMENSION + 1) X_MIN, X_MAX = 0.0, 1.0 DOM_X = HypercubeND([(X_MIN, X_MAX)] * DIMENSION, is_main_domain=True) FINAL_TIME = 0.2 DOM_T = (0.0, FINAL_TIME) SAMPLER = TensorizedSampler( [UniformTimeSampler(DOM_T), DomainSampler(DOM_X)], model_type="t_x", bc=False, ic=True, ) def f_init(x: jnp.ndarray) -> jnp.ndarray: """Initial condition: - a Gaussian centered in the middle of the domain, - and its initial time derivative, which is zero. Both outputs are required. """ x1, x2 = x[0:1], x[1:2] x1_mid, x2_mid = (X_MIN + X_MAX) / 2, (X_MIN + X_MAX) / 2 r2 = (x1 - x1_mid) ** 2 + (x2 - x2_mid) ** 2 return jnp.concatenate([jnp.exp(-75 * r2), jnp.zeros_like(r2)], axis=-1) def pre_processing(t: jnp.ndarray, x: jnp.ndarray) -> jnp.ndarray: """Pre-processing function that adds the squared distance to the center of the domain as an additional input feature.""" x1, x2 = x[0:1], x[1:2] x1_mid, x2_mid = (X_MIN + X_MAX) / 2, (X_MIN + X_MAX) / 2 r2 = (x1 - x1_mid) ** 2 + (x2 - x2_mid) ** 2 return jnp.concatenate([t, x, r2], axis=-1) # %% create and train the PINN # create the model model = WaveND( main_domain=DOM_X, time_domain=DOM_T, bc="strong", ic="weak", f_ic_rhs=f_init, ) # create the approximation space key = jax.random.PRNGKey(0) nn = MLP( in_size=DIMENSION + 2, out_size=1, hidden_sizes=[16] * (DIMENSION + 1), key=key, embedding="periodic", embedding_axes=[1, 2], periods=[1.0, 1.0], activation="sine", ) space = ApproximationSpace( {"x": DIMENSION, "t": 1}, [(nn, "scalar", None)], model_type="t_x", pre_processing=pre_processing, ) weights = {"interior": [1.0], "ic interior": [100.0, 100.0]} pinn = Projector(model, space, SAMPLER, matrix_regulazition=1e-6, weights=weights) start = timeit.default_timer() key, pinn = pinn.project(key, space, N_EPOCHS, N_COLLOC) end = timeit.default_timer() # %% exploit the results if DIMENSION <= 2: plot_abstract_approx_space( pinn.space, DOM_X, time_domain=DOM_T, time_values=(DOM_T[0], DOM_T[1] / 2, DOM_T[1]), loss=pinn.losses, title="learning sol of 2D wave equation with TemporalPinns", draw_contours=True, ) plt.show() # %%