r"""Solves a 2D Poisson PDE with Dirichlet boundary conditions using PINNs. .. math:: -\delta u & = f in \Omega where :math:`x = (x_1, x_2) \in \Omega = (-1, 1) \times (-1, 1)` and :math:`f` such that :math:`u(x_1, x_2) = \sin(\pi x_1) \sin(\pi x_2)`, and :math:`g = 0`. The boundary conditions are Dirichlet conditions enforced strongly. The neural network is a simple MLP (Multilayer Perceptron). The domain is decomposed into 4 overlapping subdomains, one per corner of the square, using a tensor product of the 1D cosine window functions used in the 1D FBPINN example along each axis. The optimization is done using a classical PINN with Natural Gradient Descent, then with a FBPINN Natural Gradient Descent. """ # %% 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.approximation_spaces.finite_basis_approximation_spaces import ( FiniteBasisApproximationSpace, make_cartesian_window_function, ) 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.laplacians import LaplacianDirichletND from scimba_jax.plots.plots_nd import ( plot_abstract_approx_spaces, ) N_COLLOC = 1000 N_EPOCHS = 50 def f_rhs(xy: jnp.ndarray) -> jnp.ndarray: x, y = xy[0:1], xy[1:2] return 2 * jnp.pi**2 * 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) * (1.0 - x) * (y + 1.0) * (1.0 - y) domain_x = (-1.0, 1.0) domain_y = (-1.0, 1.0) dx = Square2D([domain_x, domain_y], is_main_domain=True) sampler = TensorizedSampler([DomainSampler(dx)], model_type="x", bc=False) # %% print("\n\n") print("@@@@@@@@@@@@@@@ training a classic PINN @@@@@@@@@@@@@@@@@@@@@@") key = jax.random.PRNGKey(0) nn = MLP(in_size=2, out_size=1, hidden_sizes=[32, 32], key=key) space = ApproximationSpace( {"x": 2}, [(nn, "scalar", None)], model_type="x", post_processing=post_processing ) model = LaplacianDirichletND( dx, lambda *args: f_rhs(*args), model_type="x", bc="strong" ) pinn = Projector(model, space, sampler, linesearch="armijo") start = timeit.default_timer() key, pinn = pinn.project(key, space, N_EPOCHS, N_COLLOC) end = timeit.default_timer() print(f"time for {N_EPOCHS} epochs: {end - start:.2f} seconds") # %% print("\n\n") print("@@@@@@@@@@@@@@@ training an FBPINN @@@@@@@@@@@@@@@@@@@@@@") key = jax.random.PRNGKey(0) # 3 subdomains per axis, tensorized into 9 subdomains, one per corner n_subdomains_per_axis = 3 overlap = 1.5 window_function, n_subdomains = make_cartesian_window_function( [domain_x, domain_y], n_subdomains_per_axis, overlap ) if plot_window_functions := True: n_grid = 100 xs = jnp.linspace(*domain_x, n_grid) ys = jnp.linspace(*domain_y, n_grid) xx, yy = jnp.meshgrid(xs, ys, indexing="ij") pts = jnp.stack([xx.ravel(), yy.ravel()], axis=-1) fig, axes = plt.subplots(1, n_subdomains + 1, figsize=(4 * (n_subdomains + 1), 4)) sum_windows = jnp.zeros(pts.shape[0]) for i in range(n_subdomains): w = jax.vmap(lambda xy: window_function(xy, i))(pts) sum_windows += w axes[i].contourf(xx, yy, w.reshape(n_grid, n_grid), levels=20) axes[i].set_title(f"window {i}") axes[i].set_aspect("equal") im = axes[-1].contourf(xx, yy, sum_windows.reshape(n_grid, n_grid), levels=20) plt.colorbar(im, ax=axes[-1]) axes[-1].set_title("sum of windows") axes[-1].set_aspect("equal") plt.tight_layout() plt.show() dx = Square2D([domain_x, domain_y], is_main_domain=True) sampler_x = DomainSampler(dx) sampler = TensorizedSampler([sampler_x], model_type="x", bc=False) key = jax.random.PRNGKey(0) nn = [ MLP(in_size=2, out_size=1, hidden_sizes=[16, 16], key=key) for _ in range(n_subdomains) ] space2 = FiniteBasisApproximationSpace( {"x": 2}, [(nn, "scalar", None)], window_function=window_function, model_type="x", post_processing=post_processing, ) model = LaplacianDirichletND( dx, lambda *args: f_rhs(*args), model_type="x", bc="strong" ) key, sample_dict = sampler.sample(key, N_COLLOC) pinn2 = Projector( model, space2, sampler, linesearch="armijo", block_diagonal_preconditioning=True, n_subdomains=n_subdomains, truncate_jacobian_svd=False, truncate_jacobian_svd_threshold=0.02, ) start = timeit.default_timer() key, pinn2 = pinn2.project(key, space2, N_EPOCHS, N_COLLOC) end = timeit.default_timer() print(f"time for {N_EPOCHS} epochs: {end - start:.2f} seconds") plot_abstract_approx_spaces( (pinn.space, pinn2.space), dx, loss=(pinn.losses, pinn2.losses), residual=(pinn.model, pinn2.model), error=exact_sol, title="learning sol of 2D laplacian", titles=("with PINN", "with FBPINN"), draw_contours=True, n_drawn_contours=20, ) plt.show() pinn2.plot_gram_matrix() # %%