r"""Solves a 3D Poisson PDE with Dirichlet boundary conditions using PINNs. .. math:: -\delta u & = f in \Omega where :math:`x = (x_1, x_2, x_3) \in \Omega` is the unit ball and :math:`f` such that :math:`u(x_1, x_2, x_3) = \sin(\pi (x_1^2 + x_2^2 + x_3^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 7 overlapping subdomains: 1 central cube plus 6 curved caps, one per cube face, each radially projected onto the sphere -- the 3D analogue of the 2D disk example's 5-cell O-grid (1 central square + 4 curved petals). Unlike the disk example, no macro-mesh is built (GMSH has no 3D hex generator in this repository): ``make_ball_window_function`` writes each subdomain's mapping directly, since a window function only ever needs that mapping's value and Jacobian at one reference point. 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 import numpy as np from scimba_jax.domains.meshless_domains.domains_nd import BallND 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_ball_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._utils.eval_utilities import eval_on_np_tensors N_COLLOC = 2000 N_EPOCHS_PINN = 150 N_EPOCHS_FBPINN = 300 def f_rhs(xyz: jnp.ndarray) -> jnp.ndarray: x, y, z = xyz[0:1], xyz[1:2], xyz[2:3] q = x**2 + y**2 + z**2 return 4 * jnp.pi**2 * q * jnp.sin(jnp.pi * q) - 6 * jnp.pi * jnp.cos(jnp.pi * q) def exact_sol(xyz: jnp.ndarray) -> jnp.ndarray: x, y, z = xyz[:, 0:1], xyz[:, 1:2], xyz[:, 2:3] return jnp.sin(jnp.pi * (x**2 + y**2 + z**2)) def post_processing(approx: jnp.ndarray, xyz: jnp.ndarray) -> jnp.ndarray: x, y, z = xyz[0:1], xyz[1:2], xyz[2:3] return approx * (1.0 - x**2 - y**2 - z**2) def eval_space(space, xyz: np.ndarray, solution=None) -> dict: """Evaluate an approximation space (and optionally its error) on raw points. No 3D plotter exists in the library (``eval_on_np_tensors`` is what ``plot_abstract_approx_space(s)`` itself calls for 1D/2D), so this is the thinnest possible wrapper letting the custom plots below reuse it. """ n = xyz.shape[0] empty = np.zeros((n, 0)) kwargs = {} if solution is None else {"error": solution} return eval_on_np_tensors(space, empty, xyz, empty, empty, {}, 0, **kwargs) dx = BallND([0.0, 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=3, out_size=1, hidden_sizes=[24] * 3, key=key) space = ApproximationSpace( {"x": 3}, [(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) # 1 central cube + 6 curved caps (one per cube face) = 7 macro-cells, # one subdomain each -- the 3D analogue of the disk's 5-cell O-grid. window_function, n_subdomains, all_windows = make_ball_window_function( radius=1.0, inner=0.5, overlap=1.5 ) if plot_window_functions := True: n_grid_w = 100 for ax0, ax1, slice_title in [(0, 1, "z=0"), (0, 2, "y=0"), (1, 2, "x=0")]: u = jnp.linspace(-1.0, 1.0, n_grid_w) v = jnp.linspace(-1.0, 1.0, n_grid_w) uu, vv = jnp.meshgrid(u, v, indexing="ij") pts_w = ( jnp.zeros((uu.size, 3)) .at[:, ax0] .set(uu.ravel()) .at[:, ax1] .set(vv.ravel()) ) outside = np.asarray(dx.is_outside(pts_w)[:, 0]) fig = plt.figure(figsize=(20, 9)) sum_windows = jnp.zeros(pts_w.shape[0]) for i in range(n_subdomains): w = jax.vmap(lambda xyz: window_function(xyz, i))(pts_w) sum_windows += w z = np.where(outside, np.nan, np.asarray(w)).reshape(n_grid_w, n_grid_w) ax = fig.add_subplot(2, 4, i + 1) im = ax.contourf(uu, vv, z, levels=20) plt.colorbar(im, ax=ax, shrink=0.6) ax.set_aspect("equal") ax.set_title(f"window {i}") z = np.where(outside, np.nan, np.asarray(sum_windows)).reshape( n_grid_w, n_grid_w ) ax = fig.add_subplot(2, 4, n_subdomains + 1) im = ax.contourf(uu, vv, z, levels=20) plt.colorbar(im, ax=ax, shrink=0.6) ax.set_aspect("equal") ax.set_title("sum of windows") plt.suptitle(f"window functions ({slice_title})") plt.tight_layout() plt.show() dx = BallND([0.0, 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=3, out_size=1, hidden_sizes=[16] * 3, key=key) for _ in range(n_subdomains) ] space2 = FiniteBasisApproximationSpace( {"x": 3}, [(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") # %% # No 3D plotter exists in the library (plot_abstract_approx_space(s) only # handles 1D/2D): improvise with three orthogonal slices. print("\n\n") print("@@@@@@@@@@@@@@@ plotting results: slices @@@@@@@@@@@@@@@@@@@@@@") n_grid = 100 fig, axs = plt.subplots(3, 3, figsize=(13, 12)) for row, (ax0, ax1, slice_title) in enumerate( [(0, 1, "z=0"), (0, 2, "y=0"), (1, 2, "x=0")] ): u = jnp.linspace(-1.0, 1.0, n_grid) v = jnp.linspace(-1.0, 1.0, n_grid) uu, vv = jnp.meshgrid(u, v, indexing="ij") pts = jnp.zeros((uu.size, 3)).at[:, ax0].set(uu.ravel()).at[:, ax1].set(vv.ravel()) outside = np.asarray(dx.is_outside(pts)[:, 0]) exact = np.asarray(exact_sol(pts))[:, 0] err_pinn = eval_space(pinn.space, np.asarray(pts), solution=exact_sol)["error"] err_fbpinn = eval_space(pinn2.space, np.asarray(pts), solution=exact_sol)["error"] for col, (field, title) in enumerate( [(exact, "exact"), (err_pinn, "PINN error"), (err_fbpinn, "FBPINN error")] ): z = np.where(outside, np.nan, field).reshape(n_grid, n_grid) im = axs[row, col].contourf(uu, vv, z, levels=20) plt.colorbar(im, ax=axs[row, col]) axs[row, col].set_aspect("equal") axs[row, col].set_title(f"{title} ({slice_title})") plt.tight_layout() plt.show() print("\n\n") print("@@@@@@@@@@@@@@@ plotting results: convergence @@@@@@@@@@@@@@@@@@@@@@") fig, (ax_pinn, ax_fbpinn) = plt.subplots(1, 2, figsize=(11, 4)) pinn.losses.plot(ax_pinn) ax_pinn.set_title("PINN losses") pinn2.losses.plot(ax_fbpinn) ax_fbpinn.set_title("FBPINN losses") plt.tight_layout() plt.show() pinn2.plot_gram_matrix() # %%