r"""Solves a 2D parametric Helmholtz PDE with a periodic embedding. .. math:: -\delta u + k^2 u & = f in \Omega where :math:`x = (x_1, x_2) \in \Omega = (-1, 1) \times (-1, 1)`, :math:`\mu \in (0, 1)` and :math:`f` such that .. math:: u(x_1, x_2, \mu) = \cos(k \pi x_1) \cos(k \pi x_2) e^{-\mu^2}. The solution is 2-periodic in both directions, so the boundary conditions are enforced **strongly** by a *periodic embedding*: instead of feeding :math:`(x_1, x_2)` to the network, the embedding feeds :math:`(\cos(2\pi x_1 / T), \sin(2\pi x_1 / T), \ldots)`, which is periodic by construction. Whatever the network does downstream, the approximation is then exactly periodic — no boundary residual is needed, and the loss has a single interior term. The parameter :math:`\mu` is passed to the network unchanged: only the two spatial axes are embedded (``embedding_axes=[0, 1]``). The neural network is a simple MLP (Multilayer Perceptron), trained with the Energy Natural Gradient optimizer. The goal of this example is to show how to use feature enrichment in PINNs. Other embeddings exist (``"fourier"``, ``"neumann_square"``), and a comparison of several of them lives in ``examples/how_to_jax/pinns/helmholtz_2d_square_with_features.py``. """ import timeit import jax import jax.numpy as jnp import matplotlib.pyplot as plt from scimba_jax.domains.meshless_domains.base import VolumetricDomain from scimba_jax.domains.meshless_domains.domains_2d import Square2D 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_parameters import ( UniformParametricSampler, ) from scimba_jax.nonlinear_approximation.model_class.funcparam_vectorial import ( ParamScalarFunction, ) 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 ( PHYSICAL_RESIDUALS_TYPE, AbstractPhysicalModel, ) from scimba_jax.physical_models.abstract_residuals import ( NDARRAYS_FUNC_TYPE, PARAM_FUNC_TYPE, InteriorResidual, ) from scimba_jax.plots.plots_nd import plot_abstract_approx_space N_COLLOC = 2000 N_EPOCHS = 150 BOUNDS = [(-1.0, 1.0), (-1.0, 1.0)] DOM_X = Square2D(BOUNDS, is_main_domain=True) DOM_MU = [(0.0, 1.0)] # the domain is 2-wide in both directions, and the exact solution has that period PERIODS = [2.0, 2.0] # the exact solution holds a single spatial mode, so one feature per axis is # enough: measured at 150 epochs, 3 features buy nothing (relative L2 error # 2.6e-04 against 2.8e-04) for a longer training N_PERIODIC_FEATURES = 1 FREQUENCY = 1 assert FREQUENCY > 0, "FREQUENCY must be positive" assert type(FREQUENCY) in [int], "FREQUENCY must be an integer" def f_rhs(xy: jnp.ndarray, mu: jnp.ndarray) -> jnp.ndarray: x, y = xy[0:1], xy[1:2] μ = mu[0:1] k = FREQUENCY return ( (2 * jnp.pi**2 + 1) * k**2 * jnp.cos(k * jnp.pi * x) * jnp.cos(k * jnp.pi * y) * jnp.exp(-(μ**2)) ) def u_exact(x, y, μ): return ( jnp.cos(FREQUENCY * jnp.pi * x) * jnp.cos(FREQUENCY * jnp.pi * y) * jnp.exp(-(μ**2)) ) def exact_sol(xy: jnp.ndarray, mu: jnp.ndarray) -> jnp.ndarray: x, y = xy[:, 0:1], xy[:, 1:2] μ = mu[:, 0:1] return u_exact(x, y, μ) class HelmholtzResidual(InteriorResidual): def __init__( self, domain: VolumetricDomain, f_rhs: NDARRAYS_FUNC_TYPE | None = None, ): super().__init__(domain=domain, size=1, model_type="x_mu", f_rhs=f_rhs) def construct_residual(self, *vars: PARAM_FUNC_TYPE) -> PARAM_FUNC_TYPE: rho = vars[0] assert isinstance(rho, ParamScalarFunction) lap = rho.laplacian("x") return -lap + FREQUENCY**2 * rho class HelmholtzND(AbstractPhysicalModel): """A n D Helmholtz equation, with the boundary conditions enforced strongly.""" def __init__( self, main_domain: VolumetricDomain, f_rhs: NDARRAYS_FUNC_TYPE | None = None, ): super().__init__(main_domain=main_domain) self.physical_residuals: PHYSICAL_RESIDUALS_TYPE = { self.main_domain.get_label(): HelmholtzResidual( domain=main_domain, f_rhs=f_rhs, ), } key = jax.random.PRNGKey(0) # bc=False: the periodic embedding enforces the boundary conditions, so the # model carries no boundary residual and no boundary points are needed sampler = TensorizedSampler( [DomainSampler(DOM_X), UniformParametricSampler(DOM_MU)], bc=False ) nn = MLP( in_size=2 + 1, out_size=1, hidden_sizes=[16, 16], key=key, embedding="periodic", periods=PERIODS, embedding_axes=[0, 1], # only the spatial axes; mu goes through unchanged n_periodic_features=N_PERIODIC_FEATURES, ) space = ApproximationSpace({"x": 2}, [(nn, "scalar", None)], model_type="x_mu") model = HelmholtzND(DOM_X, f_rhs) print("\n\n") print("@@@@@@@@@@@@@@@ train a PINN with periodic embedding @@@@@@@@@@@@@@@@@") pinn = Projector(model, space, sampler) key, sample_dict = sampler.sample(key, N_COLLOC) print("initial loss: ", pinn.evaluate_loss(space, sample_dict)) start = timeit.default_timer() key, pinn = pinn.project(key, space, N_EPOCHS, N_COLLOC) end = timeit.default_timer() print("best loss: ", pinn.best_loss) print("time for %d epochs: " % N_EPOCHS, end - start) plot_abstract_approx_space( pinn.space, # the approximation space DOM_X, # the spatial domain DOM_MU, # the parameters domain loss=pinn.losses, # for plot of the loss: the losses residual=pinn.model, # for plot of the residual: the pde solution=exact_sol, # for plot of the exact sol: sol error=exact_sol, # for plot of the error with respect to a func: the func derivatives=["ux", "uy"], # a list of strings for the derivatives to plot # for plots on linear cuts of dim d-1, a tuple (point, direction) cuts=[([0.0, 0.0], [1.0, 1.0])], parameters_values="mean", # the value of mu at which the 2D maps are drawn draw_contours=True, n_drawn_contours=20, title="2D parametric Helmholtz with a periodic embedding", ) plt.show()