r"""Solves a 2D reaction-diffusion PDE with adaptive ENG matrix regularization. .. math:: -\delta u + u & = f in \Omega where :math:`x = (x_1, x_2) \in \Omega = (0, 1) \times (0, 1)` and :math:`f` such that :math:`u(x_1, x_2) = \sin(\pi x_1) \sin(\pi x_2)`, with homogeneous Dirichlet boundary conditions enforced strongly. The neural network is a simple MLP (Multilayer Perceptron), trained with the Energy Natural Gradient optimizer, an Armijo linesearch and an *adaptive* matrix regularization. The Energy Natural Gradient optimizer damps its Gram matrix by ``matrix_regularization``. Kept fixed, that value has to be hand-tuned: too large and the natural-gradient step degenerates into a tiny, slow gradient step; too small and the (near-singular, early in training) Gram matrix produces a huge, unstable step that the linesearch keeps rejecting. With ``adaptive_matrix_regularization=True`` the damping is instead adjusted every epoch à la Levenberg-Marquardt, keyed on the outcome of the linesearch: relaxed when the full step is accepted outright (the Gram matrix is trusted), tightened when the search fails to decrease the loss. The initial value below is therefore only a starting point, not a tuned constant. This is the JAX port of the ``adaptive_matrix_regularization`` option of the torch ``EnergyNaturalGradientPreconditioner``; the torch example that turns it on is ``examples/examples_torch/pinns/stationary_pdes/elliptic_pdes/ reaction_diffusion_2d_square_parametric.py``. """ import timeit import jax import jax.numpy as jnp import matplotlib.pyplot as plt 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.networks.mlp import MLP from scimba_jax.nonlinear_approximation.numerical_solvers.projectors import Projector from scimba_jax.physical_models.elliptic_pde.general_elliptic import GeneralElliptic from scimba_jax.plots.plots_nd import plot_abstract_approx_space N_COLLOC = 2000 N_EPOCHS = 100 # only a starting point: the adaptation below grows or shrinks it every epoch INITIAL_MATRIX_REGULARIZATION = 1e-6 key = jax.random.PRNGKey(0) def f_rhs(xy: jnp.ndarray) -> jnp.ndarray: x, y = xy[0:1], xy[1:2] return (2 * jnp.pi**2 + 1.0) * jnp.sin(jnp.pi * x) * jnp.sin(jnp.pi * y) def exact_sol(xy: jnp.ndarray) -> jnp.ndarray: x, y = xy[:, 0:1], xy[:, 1:2] return jnp.sin(jnp.pi * x) * jnp.sin(jnp.pi * y) def post_processing(approx: jnp.ndarray, xy: jnp.ndarray) -> jnp.ndarray: x, y = xy[0:1], xy[1:2] return approx * x * (1.0 - x) * y * (1.0 - y) domain_x = [(0.0, 1.0), (0.0, 1.0)] dx = Square2D(domain_x, is_main_domain=True) sampler = TensorizedSampler([DomainSampler(dx)], bc=False) # -div(grad(u)) + u = f, i.e. A = I (default), b = None, c = 1 model = GeneralElliptic( dx, model_type="x", f_rhs=lambda *args: f_rhs(*args), bc="strong", c=1.0 ) nn = MLP(in_size=2, out_size=1, hidden_sizes=[16, 16], key=key) space = ApproximationSpace( {"x": 2}, [(nn, "scalar", None)], model_type="x", post_processing=post_processing, ) print("\n\n") print("@@@@@@@@@@@@@@@ train with ENG, armijo, adaptive reg. @@@@@@@@@@@@@@@@@@@@@") pinn = Projector( model, space, sampler, linesearch="armijo", matrix_regularization=INITIAL_MATRIX_REGULARIZATION, adaptive_matrix_regularization=True, ) 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) print( "matrix_regularization: %.2e (initial) -> %.2e (final)" % (INITIAL_MATRIX_REGULARIZATION, float(pinn.optimizer.matrix_regularization)) ) plot_abstract_approx_space( pinn.space, # the approximation space dx, # the spatial domain loss=pinn.losses, # the losses residual=pinn.model, # the pde error=exact_sol, # the error with respect to the exact sol draw_contours=True, n_drawn_contours=20, title=( "2D reaction-diffusion with ENG, armijo linesearch " "and adaptive matrix regularization" ), ) plt.show()