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` is the unit disk and :math:`f` such that :math:`u(x_1, x_2) = \sin(\pi (x_1^2 + x_2^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 5 overlapping subdomains using the O-grid macro-mesh of the disk (one central square, four curved petals), with window functions built from the macro-mesh's own cells, as in gmsh_subdomains_exploration.py. 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 Disk2D from scimba_jax.mapping.macro_mesh import macro_mesh_ogrid_disk 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_mesh_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_PINN = 150 N_EPOCHS_FBPINN = 300 def f_rhs(xy: jnp.ndarray) -> jnp.ndarray: x, y = xy[0:1], xy[1:2] q = x**2 + y**2 return 4 * jnp.pi**2 * q * jnp.sin(jnp.pi * q) - 4 * jnp.pi * jnp.cos(jnp.pi * q) def exact_sol(xy: jnp.ndarray) -> jnp.ndarray: x, y = xy[:, 0:1], xy[:, 1:2] return jnp.sin(jnp.pi * (x**2 + y**2)) def post_processing(approx: jnp.ndarray, xy: jnp.ndarray) -> jnp.ndarray: x, y = xy[0:1], xy[1:2] return approx * (1.0 - x**2 - y**2) dx = Disk2D([0.0, 0.0], 1.0, 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_PINN, N_COLLOC) end = timeit.default_timer() print(f"time for {N_EPOCHS_PINN} epochs: {end - start:.2f} seconds") # %% print("\n\n") print("@@@@@@@@@@@@@@@ training an FBPINN @@@@@@@@@@@@@@@@@@@@@@") key = jax.random.PRNGKey(0) # O-grid disk, n=1: 1 central square + 4 curved petals = 5 macro-cells, # one subdomain each. mesh = macro_mesh_ogrid_disk(radius=1.0, inner=0.5, n=1) window_function, n_subdomains, _ = make_mesh_window_function(mesh, overlap=1.5) if plot_window_functions := True: n_grid = 100 xs = jnp.linspace(-1.0, 1.0, n_grid) ys = jnp.linspace(-1.0, 1.0, n_grid) xx, yy = jnp.meshgrid(xs, ys, indexing="ij") pts = jnp.stack([xx.ravel(), yy.ravel()], axis=-1) outside = dx.is_outside(pts)[:, 0] 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 w = jnp.where(outside, jnp.nan, 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") sum_windows = jnp.where(outside, jnp.nan, sum_windows) 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 = Disk2D([0.0, 0.0], 1.0, 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) # block_diagonal_preconditioning is not used here # because the Gram matrix is no longer block-diagonal pinn2 = Projector(model, space2, sampler, linesearch="armijo") start = timeit.default_timer() key, pinn2 = pinn2.project(key, space2, N_EPOCHS_FBPINN, N_COLLOC) end = timeit.default_timer() print(f"time for {N_EPOCHS_FBPINN} 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 on the disk", titles=("with PINN", "with FBPINN"), draw_contours=True, n_drawn_contours=20, ) plt.show() pinn2.plot_gram_matrix() # %%