r"""Solves a 3D Poisson PDE written directly in axisymmetric coordinates. .. math:: -\left(\frac{\partial^2 u}{\partial r^2} + \frac{1}{r} \frac{\partial u}{\partial r} + \frac{\partial^2 u}{\partial z^2}\right) & = f \quad \text{in } \Omega = (r_{\min}, r_{\max}) \times (z_{\min}, z_{\max}) \\ u & = g \quad \text{on } \partial \Omega This is the companion example to ``poisson_2d_cylindrical.py``: instead of a 2D problem in the polar plane :math:`(r, \theta)`, this is a 3D problem with no dependence on the azimuthal angle :math:`\theta` (axisymmetric), so the unknown only depends on :math:`(r, z)`, with :math:`r` the distance to the symmetry axis and :math:`z` the coordinate along that axis. As in the polar case, the physical domain, with coordinates :math:`(r, z)`, is generated from a reference unit square through a :class:`~scimba_jax.domains.domain_mapping.DomainMapping` -- the same mechanism as the tokamak (Grad-Shafranov) example. The neural network takes :math:`(r, z)` directly as input, and the Laplace operator is written in its axisymmetric form, with an explicit :math:`1/r` factor (and no :math:`1/r^2` term, since there is no :math:`\theta`-dependence). An exact solution :math:`u(r, z) = (1 - r^2) \sin(\pi z)` is manufactured -- smooth and regular on the symmetry axis :math:`r = 0` (it only depends on :math:`r` through :math:`r^2`, so :math:`u_r / r` stays finite there) -- so that the right-hand side :math:`f` and the (non-homogeneous) Dirichlet data :math:`g` can be computed analytically, and errors can be measured. Boundary conditions are enforced weakly. The neural network is a simple MLP, trained with the default (energy natural gradient) optimizer. """ import timeit import jax import jax.numpy as jnp import matplotlib.pyplot as plt from scimba_jax.domains.domain_mapping import DomainMapping 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.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.curvilinear_poisson import ( AxisymmetricPoissonDirichletND, ) from scimba_jax.plots.plots_nd import plot_abstract_approx_space # ── Parameters ───────────────────────────────────────────────────────────────── R_MIN, R_MAX = 0.0, 1.0 Z_MIN, Z_MAX = 0.0, 1.0 N_COLLOC = 2000 N_BC_COLLOC = 2000 N_EPOCHS = 50 HIDDEN = [16] * 3 key = jax.random.PRNGKey(0) # ── Manufactured solution and source term ────────────────────────────────────── # u(r, z) = (1 - r**2) * sin(pi * z) # f = -Delta u, with Delta the axisymmetric Laplacian. def exact_sol_pointwise(x: jnp.ndarray) -> jnp.ndarray: r, z = x[0:1], x[1:2] return (1.0 - r**2) * jnp.sin(jnp.pi * z) def exact_sol(x: jnp.ndarray) -> jnp.ndarray: """Batched exact solution, for plotting.""" r, z = x[:, 0:1], x[:, 1:2] return (1.0 - r**2) * jnp.sin(jnp.pi * z) def f_rhs(x: jnp.ndarray) -> jnp.ndarray: r, z = x[0:1], x[1:2] return jnp.sin(jnp.pi * z) * (4.0 + (1.0 - r**2) * jnp.pi**2) def dirichlet_bc(x: jnp.ndarray, n: jnp.ndarray) -> jnp.ndarray: return exact_sol_pointwise(x) # ── Domain: axisymmetric (r, z) coordinates generated via a mapping ─────────── # A reference unit square is mapped to the physical (r, z) rectangle, exactly # as the tokamak example maps a reference unit disk to the physical (R, Z) # tokamak cross-section. mapping = DomainMapping.rectangle( origin=jnp.array([R_MIN, Z_MIN]), x_dir=jnp.array([R_MAX - R_MIN, 0.0]), y_dir=jnp.array([0.0, Z_MAX - Z_MIN]), ) domain = Square2D([(0.0, 1.0), (0.0, 1.0)], is_main_domain=True) domain.set_mapping(mapping, bounds_postmap=[(R_MIN, R_MAX), (Z_MIN, Z_MAX)]) sampler = TensorizedSampler([DomainSampler(domain)], bc=True) # ── Network & space ──────────────────────────────────────────────────────────── key, subkey = jax.random.split(key) nn = MLP(in_size=2, out_size=1, hidden_sizes=HIDDEN, key=subkey) space = ApproximationSpace({"x": 2}, [(nn, "scalar", None)], model_type="x") model = AxisymmetricPoissonDirichletND( domain, f_rhs=lambda *args: f_rhs(*args), bc="weak", f_bc_rhs=lambda *args: dirichlet_bc(*args), model_type="x", ) # ── Training ──────────────────────────────────────────────────────────────────── weights_dict = {"interior": [1.0], "boundary": [40.0]} pinn = Projector(model, space, sampler, weights=weights_dict) key, sample_dict = sampler.sample(key, N_COLLOC, N_BC_COLLOC) print("initial loss: ", pinn.evaluate_loss(space, sample_dict)) t0 = timeit.default_timer() key, pinn = pinn.project(key, space, N_EPOCHS, N_COLLOC, N_BC_COLLOC) print("best loss: ", pinn.best_loss) print("time for %d epochs: " % N_EPOCHS, timeit.default_timer() - t0) # ── Plot ───────────────────────────────────────────────────────────────────── # horizontal axis: r, vertical axis: z plot_abstract_approx_space( pinn.space, domain, loss=pinn.losses, residual=pinn.model, error=exact_sol, draw_contours=True, n_drawn_contours=20, title="3D Poisson equation in axisymmetric (r, z) coordinates", ) plt.show()