r"""Solves a 1D Poisson PDE with Dirichlet boundary conditions using PINNs. .. math:: -\delta u & = f in \Omega where :math:`x \in \Omega = (0, 1)` and :math:`f` such that :math:`u(x) = \sin(\pi x)`, and :math:`g = 0`. The boundary conditions are Dirichlet conditions enforced strongly. The neural network is a simple MLP (Multilayer Perceptron). 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_1d import Segment1D 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.integration.monte_carlo_parameters import ( # UniformParametricSampler, # ) 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_space, ) N_COLLOC = 1000 N_EPOCHS = 50 def f_rhs(x: jnp.ndarray): return jnp.pi**2 * jnp.sin(jnp.pi * x) def f_bc(x: jnp.ndarray, n: jnp.ndarray): return x * 0 def exact_sol(x: jnp.ndarray): return jnp.sin(jnp.pi * x) def post_processing(approx: jnp.ndarray, x: jnp.ndarray): return approx * x * (1 - x) domain_x = (0, 1) dx = Segment1D(domain_x, 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=1, out_size=1, hidden_sizes=[32, 32], key=key) space = ApproximationSpace( {"x": 1}, [(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) n_subdomains = 4 overlap = 1.5 window_function, n_subdomains = make_cartesian_window_function( [domain_x], n_subdomains, overlap ) if plot_window_functions := False: x = jnp.linspace(0, 1, 100) for i in range(n_subdomains): plt.plot(x, window_function(x, i), label=f"window {i}") sum_windows = 0 for i in range(n_subdomains): sum_windows += window_function(x, i) plt.plot(x, sum_windows, label="sum of windows", linestyle="--") plt.legend() dx = Segment1D(domain_x, is_main_domain=True) sampler_x = DomainSampler(dx) sampler = TensorizedSampler([sampler_x], model_type="x", bc=False) # print("\n\n") # print("@@@@@@@@@@@@@@@ define an FBPINN @@@@@@@@@@@@@@@@@@@@@@") key = jax.random.PRNGKey(0) nn = [ MLP(in_size=1, out_size=1, hidden_sizes=[16, 16], key=key) for _ in range(n_subdomains) ] space = FiniteBasisApproximationSpace( {"x": 1}, [(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) pinn = Projector( model, space, sampler, linesearch="armijo", block_diagonal_preconditioning=True, n_subdomains=n_subdomains, truncate_jacobian_svd=False, truncate_jacobian_svd_threshold=0.02, ) print("\n\n") print("@@@@@@@@@@@@@@@ Train an FBPINN with ENG @@@@@@@@@@@@@@@@@@@@@@") 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") plot_abstract_approx_space( pinn.space, dx, loss=pinn.losses, residual=pinn.model, solution=exact_sol, error=exact_sol, title="learning sol of 1D laplacian with an FBPINN", ) plt.show() pinn.plot_gram_matrix() # %%