r"""Solves a 4D Poisson PDE with Dirichlet boundary conditions using PINNs. .. math:: -\Delta u & = f \quad \text{in } \Omega where :math:`x = (x_1, x_2, x_3, x_4) \in \Omega = (-1, 1)^4` and :math:`f` is chosen such that the exact solution is: .. math:: u(x) = \prod_{i=1}^{4} \sin(\pi x_i) which gives :math:`f(x) = 4\pi^2 \prod_{i=1}^{4} \sin(\pi x_i)` and homogeneous Dirichlet boundary conditions. The boundary conditions are enforced strongly via a post-processing that vanishes on :math:`\partial \Omega` (``HypercubeND`` has no boundary domain, so weak BC is not available here). The neural network is a simple MLP (Multilayer Perceptron), trained with the Energy Natural Gradient. """ # %% import timeit import matplotlib.pyplot as plt import torch from scimba_torch.approximation_space.nn_space import NNxSpace from scimba_torch.domain.meshless_domain.domain_nd import HypercubeND from scimba_torch.integration.monte_carlo import DomainSampler, TensorizedSampler from scimba_torch.integration.monte_carlo_parameters import UniformParametricSampler from scimba_torch.neural_nets.coordinates_based_nets.mlp import GenericMLP from scimba_torch.numerical_solvers.elliptic_pde.pinns import ( NaturalGradientPinnsElliptic, ) from scimba_torch.physical_models.elliptic_pde.linear_order_2 import LinearOrder2PDE from scimba_torch.utils.scimba_tensors import LabelTensor torch.manual_seed(0) N_COLLOC = 10000 N_BC_COLLOC = 10000 N_EPOCHS = 200 dim = 4 domain_x = [(-1.0, 1.0)] * dim def f_rhs(x: LabelTensor, mu: LabelTensor) -> torch.Tensor: """RHS such that u = prod_i sin(pi * x_i).""" x_components = x.get_components() result = dim * torch.pi**2 for x_i in x_components: result = result * torch.sin(torch.pi * x_i) return result def exact_sol(x: LabelTensor, mu: LabelTensor) -> torch.Tensor: """Exact solution u = prod_i sin(pi * x_i), batched.""" result = 1.0 for x_i in x.get_components(): result = result * torch.sin(torch.pi * x_i) return result class Laplacian4DDirichletStrongFormNoParam(LinearOrder2PDE): """Plain -Delta u = f, with no dependency on a parameter mu.""" def operator(self, w, x, mu): u = w.get_components() grad_u = torch.cat(tuple(self.grad(u, x)), dim=-1) div_grad_u = tuple(self.grad(grad_u[:, 0], x))[0] for i in range(1, self.spatial_dim): div_grad_u = div_grad_u + tuple(self.grad(grad_u[:, i], x))[i] return -div_grad_u def functional_operator(self, func, x, mu, theta): grad_u = torch.func.jacrev(func, 0) hessian_u = torch.func.jacrev(grad_u, 0, chunk_size=None)(x, mu, theta) return -sum(hessian_u[..., i, i] for i in range(self.spatial_dim)) dx = HypercubeND(domain_x, is_main_domain=True) # %% print( "@@@@@@@@@@@@@@@ create a PINN with strong BC (4D Laplacian) @@@@@@@@@@@@@@@@@@@@@" ) def post_processing( inputs: torch.Tensor, x: LabelTensor, mu: LabelTensor ) -> torch.Tensor: """Enforce homogeneous Dirichlet BC strongly: multiply by prod_i (1 - x_i^2).""" factor = 1.0 for x_i in x.get_components(): factor = factor * (1.0 - x_i**2) return inputs * factor def functional_post_processing(u, x, mu, theta): """Pointwise counterpart of post_processing, for the natural gradient.""" factor = 1.0 for i in range(dim): factor = factor * (1.0 - x[i] ** 2) return u(x, mu, theta) * factor sampler = TensorizedSampler([DomainSampler(dx), UniformParametricSampler([])]) space = NNxSpace( 1, 0, GenericMLP, dx, sampler, layer_sizes=[24, 24, 24], post_processing=post_processing, ) pde = Laplacian4DDirichletStrongFormNoParam(space, spatial_dim=dim, f=f_rhs) pinn = NaturalGradientPinnsElliptic( pde, bc_type="strong", matrix_regularization=1e-4, functional_post_processing=functional_post_processing, ) start = timeit.default_timer() pinn.solve(epochs=N_EPOCHS, n_collocation=N_COLLOC) end = timeit.default_timer() print(f"best loss: {pinn.best_loss:.4e}") print(f"time for {N_EPOCHS} epochs: {end - start:.2f}s") # %% def plot_2d_proj(pinn, title, n_visu=128, val_fix=0.5): """Plot 2D cuts fixing x3=x4=val_fix, comparing prediction, exact and error.""" x1_lin = torch.linspace(-1, 1, n_visu) x2_lin = torch.linspace(-1, 1, n_visu) X1, X2 = torch.meshgrid(x1_lin, x2_lin, indexing="xy") x_flat = torch.stack([X1.reshape(-1), X2.reshape(-1)], dim=-1) # fix x3=x4=0.5 (x3=x4=0 makes sin(pi*0)=0 => exact sol identically zero) x_cut = torch.cat([x_flat, val_fix * torch.ones((x_flat.shape[0], 2))], dim=-1) x_label = LabelTensor(x_cut) mu_label = LabelTensor(torch.zeros((x_flat.shape[0], 0))) u_pred = pinn.evaluate(x_label, mu_label).w.detach() u_exact = exact_sol(x_label, mu_label) u_pred = u_pred.reshape(n_visu, n_visu) u_exact = u_exact.reshape(n_visu, n_visu) error = torch.abs(u_pred - u_exact) fig, axes = plt.subplots(1, 3, figsize=(15, 4)) im0 = axes[0].contourf(X1, X2, u_pred, levels=32, cmap="turbo") plt.colorbar(im0, ax=axes[0]) axes[0].set_title("PINN solution (x3=x4=0.5)") axes[0].set_xlabel("x1") axes[0].set_ylabel("x2") im1 = axes[1].contourf(X1, X2, u_exact, levels=32, cmap="turbo") plt.colorbar(im1, ax=axes[1]) axes[1].set_title("Exact solution (x3=x4=0.5)") axes[1].set_xlabel("x1") axes[1].set_ylabel("x2") im2 = axes[2].contourf(X1, X2, error, levels=32, cmap="inferno") plt.colorbar(im2, ax=axes[2]) axes[2].set_title(f"Absolute error (max={float(error.max()):.2e})") axes[2].set_xlabel("x1") axes[2].set_ylabel("x2") plt.suptitle(title) plt.tight_layout() plt.show() n_test = 5000 x_test = LabelTensor( torch.stack([torch.empty(n_test).uniform_(-1, 1) for _ in range(dim)], dim=-1) ) mu_test = LabelTensor(torch.zeros((n_test, 0))) u_pred_test = pinn.evaluate(x_test, mu_test).w.detach() u_exact_test = exact_sol(x_test, mu_test) l2_error = torch.sqrt(torch.mean((u_pred_test - u_exact_test) ** 2)) print(f"L2 error on {n_test} test points: {float(l2_error):.4e}") plot_2d_proj(pinn, "4D Laplacian, strong BC — cut at x3=x4=0.5")