r"""Solves a 1D heat equation with a varying diffusion coefficient using a PINN. .. math:: \partial_t u - \partial_x (k(t, x) \partial_x u) & = f in \Omega \times (0, T) \\ u & = g on \partial \Omega \times (0, T) \\ u & = u_0 on \Omega \times \{0\} where :math:`u: \Omega \times (0, T) \to \mathbb{R}` is the unknown function, :math:`\Omega = (-1, 1) \subset \mathbb{R}` is the spatial domain, :math:`(0, T) = (0, 1) \subset \mathbb{R}` is the time domain and .. math:: u(t, x) = \tanh(xt), \qquad k(t, x) = \frac{1}{\cosh(xt)}, \qquad f(t, x) = \frac{1}{\cosh^2(xt)} \left(\frac{3t^2 \tanh(xt)}{\cosh(xt)} + x\right). Dirichlet boundary conditions :math:`u(t, -1) = \tanh(-t)`, :math:`u(t, 1) = \tanh(t)` are prescribed; both coincide with the exact solution evaluated at the boundary, so ``exact_sol`` doubles as the boundary data. The equation is solved on a segment domain with weak boundary and initial conditions, reusing :class:`~scimba_jax.physical_models.temporal_pde.heat_equations.AnisotropicHeatND` (the generalization of ``HeatND`` to a variable diffusion coefficient) with energy natural gradient preconditioning, strong-wolfe linesearch and adaptive matrix regularization. """ 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.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.heat_equations import AnisotropicHeatND from scimba_jax.plots.plots_nd import plot_abstract_approx_space N_COLLOC = 2000 N_BC_COLLOC = 2000 N_IC_COLLOC = 2000 N_EPOCHS = 1000 def exact_sol(t: jnp.ndarray, x: jnp.ndarray) -> jnp.ndarray: return jnp.tanh(x * t) def diffusion_coeff(t: jnp.ndarray, x: jnp.ndarray) -> jnp.ndarray: """Diffusion coefficient k(t, x), as the 1x1 matrix anisotropic_laplacian expects.""" return (1.0 / jnp.cosh(x * t)).reshape(1, 1) def f_rhs(t: jnp.ndarray, x: jnp.ndarray) -> jnp.ndarray: return (1.0 / jnp.cosh(x * t) ** 2) * ( 3 * t**2 * jnp.tanh(x * t) / jnp.cosh(x * t) + x ) def f_init(x: jnp.ndarray) -> jnp.ndarray: t = jnp.zeros_like(x) return exact_sol(t, x) def dirichlet_bc(t: jnp.ndarray, x: jnp.ndarray, n: jnp.ndarray) -> jnp.ndarray: return exact_sol(t, x) domain_t = (0.0, 10.0) domain_x = [(-1.0, 1.0)] dx = Segment1D(domain_x[0], is_main_domain=True) sampler = TensorizedSampler( [ UniformTimeSampler(domain_t), DomainSampler(dx), ], model_type="t_x", bc=True, ic=True, ) # create the model model = AnisotropicHeatND( main_domain=dx, time_domain=domain_t, bc="weak", ic="weak", f_rhs=f_rhs, A=diffusion_coeff, f_bc_rhs=lambda *args: dirichlet_bc(*args), f_ic_rhs=lambda *args: f_init(*args), ) # create the approximation space key = jax.random.PRNGKey(0) nn = MLP(in_size=2, out_size=1, hidden_sizes=[12] * 3, key=key) space = ApproximationSpace( {"x": 1}, [(nn, "scalar", None)], model_type="t_x", ) # create the pinn, trained with energy natural gradient preconditioning pinn = Projector( model, space, sampler, linesearch="strong-wolfe", adaptive_matrix_regularization=True, ) start = timeit.default_timer() key, pinn = pinn.project(key, space, N_EPOCHS, N_COLLOC, N_BC_COLLOC, N_IC_COLLOC) end = timeit.default_timer() print("time for %d epochs: " % N_EPOCHS, end - start) plot_abstract_approx_space( pinn.space, dx, time_domain=domain_t, time_values=[0.0, 2.5, 5.0, 7.5, 10.0], loss=pinn.losses, residual=pinn.model, solution=exact_sol, error=exact_sol, title="learning sol of 1D heat equation with variable diffusion coefficient", ) plt.show()