r"""Solves the heat equation written directly in cylindrical (polar) coordinates. .. math:: \partial_t u - \left(\frac{\partial^2 u}{\partial r^2} + \frac{1}{r} \frac{\partial u}{\partial r} + \frac{1}{r^2} \frac{\partial^2 u}{\partial \theta^2}\right) & = f \quad \text{in } \Omega \times (0, T) \\ u & = g \quad \text{on } \partial \Omega \times (0, T) \\ u & = u_0 \quad \text{on } \Omega \times \{0\} where :math:`\Omega = (r_{\min}, r_{\max}) \times (\theta_{\min}, \theta_{\max})`. As in the companion elliptic example (``poisson_2d_cylindrical.py``), the physical domain -- with coordinates :math:`(r, \theta)` -- is generated from a reference unit square through a :class:`~scimba_jax.domains.domain_mapping.DomainMapping`, following the same mechanism as the tokamak (Grad-Shafranov) example. The neural network takes :math:`(t, r, \theta)` as input, and the diffusion operator is written in its cylindrical form, with explicit :math:`1/r` and :math:`1/r^2` factors. An exact solution :math:`u(t, r, \theta) = e^{-t} \sin(\alpha (r - r_{\min})) \cos(m \theta)`, with :math:`\alpha = \pi / (r_{\max} - r_{\min})`, is manufactured (via a source term :math:`f`) so that errors can be measured. Boundary and initial 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.integration.monte_carlo_time import ( UniformTimeSampler, ) from scimba_jax.nonlinear_approximation.networks.mlp import MLP from scimba_jax.nonlinear_approximation.numerical_solvers.projectors import Projector from scimba_jax.physical_models.temporal_pde.curvilinear_heat import CylindricalHeatND from scimba_jax.plots.plots_nd import plot_abstract_approx_space # ── Parameters ───────────────────────────────────────────────────────────────── R_MIN, R_MAX = 1.0, 2.0 THETA_MIN, THETA_MAX = 0.0, 2.0 * jnp.pi T_MIN, T_MAX = 0.0, 1.0 M = 2 ALPHA = jnp.pi / (R_MAX - R_MIN) N_COLLOC = 2500 N_BC_COLLOC = 2000 N_IC_COLLOC = 2000 N_EPOCHS = 100 HIDDEN = [16] * 3 key = jax.random.PRNGKey(0) # ── Manufactured solution and source term ────────────────────────────────────── # u(t, r, theta) = exp(-t) * sin(alpha * (r - r_min)) * cos(m * theta) # f = u_t - Delta u, with Delta the cylindrical Laplacian. def exact_sol_pointwise(t: jnp.ndarray, x: jnp.ndarray) -> jnp.ndarray: r, theta = x[0:1], x[1:2] return jnp.exp(-t) * jnp.sin(ALPHA * (r - R_MIN)) * jnp.cos(M * theta) def exact_sol(t: jnp.ndarray, x: jnp.ndarray) -> jnp.ndarray: """Batched exact solution, for plotting.""" r, theta = x[:, 0:1], x[:, 1:2] return jnp.exp(-t) * jnp.sin(ALPHA * (r - R_MIN)) * jnp.cos(M * theta) def f_init(x: jnp.ndarray) -> jnp.ndarray: return exact_sol_pointwise(jnp.zeros(1), x) def f_rhs(t: jnp.ndarray, x: jnp.ndarray) -> jnp.ndarray: r, theta = x[0:1], x[1:2] s = jnp.sin(ALPHA * (r - R_MIN)) c = jnp.cos(ALPHA * (r - R_MIN)) bracket = (ALPHA**2 + (M**2) / r**2 - 1.0) * s - ALPHA * c / r return jnp.exp(-t) * jnp.cos(M * theta) * bracket def dirichlet_bc(t: jnp.ndarray, x: jnp.ndarray, n: jnp.ndarray) -> jnp.ndarray: return exact_sol_pointwise(t, x) # ── Domain: cylindrical (r, theta) coordinates generated via a mapping ──────── # A reference unit square is mapped to the physical (r, theta) 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, THETA_MIN]), x_dir=jnp.array([R_MAX - R_MIN, 0.0]), y_dir=jnp.array([0.0, THETA_MAX - THETA_MIN]), ) domain_x = Square2D([(0.0, 1.0), (0.0, 1.0)], is_main_domain=True) domain_x.set_mapping(mapping, bounds_postmap=[(R_MIN, R_MAX), (THETA_MIN, THETA_MAX)]) domain_t = (T_MIN, T_MAX) sampler = TensorizedSampler( [UniformTimeSampler(domain_t), DomainSampler(domain_x)], model_type="t_x", bc=True, ic=True, ) # ── Network & space ──────────────────────────────────────────────────────────── key, subkey = jax.random.split(key) nn = MLP(in_size=3, out_size=1, hidden_sizes=HIDDEN, key=subkey) space = ApproximationSpace({"x": 2, "t": 1}, [(nn, "scalar", None)], model_type="t_x") model = CylindricalHeatND( main_domain=domain_x, time_domain=domain_t, bc="weak", f_rhs=lambda *args: f_rhs(*args), f_bc_rhs=lambda *args: dirichlet_bc(*args), ic="weak", f_ic_rhs=lambda *args: f_init(*args), ) # ── Training ──────────────────────────────────────────────────────────────────── pinn = Projector(model, space, sampler) key, sample_dict = sampler.sample(key, N_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, N_IC_COLLOC) print("best loss: ", pinn.best_loss) print("time for %d epochs: " % N_EPOCHS, timeit.default_timer() - t0) # ── Plot ───────────────────────────────────────────────────────────────────── plot_abstract_approx_space( pinn.space, domain_x, time_domain=domain_t, time_values=(T_MIN, (T_MIN + T_MAX) / 2, T_MAX), loss=pinn.losses, residual=pinn.model, error=exact_sol, title="2D heat equation in cylindrical (r, theta) coordinates", ) plt.show()