r"""Solves a 1D heat equation with varying diffusion using a discrete 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). Discrete PINN, not continuous: the diffusion coefficient sits *inside* the Laplace-type operator, not just in the source term ``f``. A discrete PINN's spatial network only ever sees ``x`` (time is handled by the outer Runge-Kutta loop), so unlike ``f`` -- already evaluated at each stage's exact time by :class:`~scimba_jax.nonlinear_approximation.numerical_solvers.discrete_pinns. DiscretePINN` itself -- a time-varying *operator* needs the PDE object rebuilt at each stage's own time. The spatial operator itself -- ``-div(k(t, .) grad u)`` at a frozen, known ``t`` -- is nothing but a steady variable-coefficient Poisson problem, so it is reused as-is from :class:`~scimba_jax.physical_models.elliptic_pde. general_elliptic.GeneralElliptic` (used elsewhere in the library for exactly this operator) rather than writing a new residual class, built with ``bc="strong"`` (no boundary residual from ``GeneralElliptic`` itself -- its own ``bc="weak"`` wiring hardcodes a ``(x, mu)`` signature for a callable ``A``, which does not match a ``mu``-less ``model_type="x"`` like this one). The Dirichlet boundary condition :math:`u(t, \pm 1) = \tanh(\pm t)` is instead added directly as a :class:`~scimba_jax.physical_models.boundary_residuals. DirichletResidual` per boundary label, exactly as every ``bc="weak"`` PDE in the library already does -- also rebuilt per stage, since the boundary data is itself time-varying. """ 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 ( # noqa: E501 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.discrete_pinns import ( DiscretePINN, ) from scimba_jax.physical_models.abstract_physical_model import AbstractPhysicalModel from scimba_jax.physical_models.boundary_residuals import DirichletResidual from scimba_jax.physical_models.elliptic_pde.general_elliptic import GeneralElliptic from scimba_jax.plots.plots_nd import plot_abstract_approx_spaces from scimba_jax.time_discrete.butcher_tableau import ( build_alexander_tableau, build_implicit_euler_tableau, ) N_COLLOC = 1000 N_BC_COLLOC = 2 N_EPOCHS_INIT = 2500 N_EPOCHS = 15 DOM_X = Segment1D((-1.0, 1.0), is_main_domain=True) DOM_T = (0.0, 10.0) SAMPLER = TensorizedSampler([DomainSampler(DOM_X)], model_type="x", bc=True) 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: return exact_sol(jnp.zeros_like(x), x) def pde_factory(t: jnp.ndarray) -> AbstractPhysicalModel: """Build the steady, variable-coefficient Poisson problem for time ``t``. ``A`` and ``f_bc_rhs`` close over ``t`` (frozen at this stage's own time); ``f_rhs`` does not need to, since ``DiscretePINN`` already re-evaluates it itself at each stage's exact time. Reused, unchanged, as either ``implicit_pde`` or ``explicit_pde`` below. Args: t: the (possibly traced) time at which to freeze the operator and the boundary data. Returns: The steady PDE ``-div(k(t, .) grad u) = f(., .)`` with Dirichlet data ``u = exact_sol(t, .)``, ready to be combined into a Runge-Kutta stage by :class:`DiscretePINN`. """ def diffusion_coeff_at_t(x: jnp.ndarray) -> jnp.ndarray: return diffusion_coeff(t, x) def dirichlet_bc_at_t(x: jnp.ndarray, n: jnp.ndarray) -> jnp.ndarray: return exact_sol(t, x) pde = GeneralElliptic( main_domain=DOM_X, model_type="x", f_rhs=f_rhs, A=diffusion_coeff_at_t, bc="strong", ) for boundary in pde.boundaries: pde.physical_residuals[boundary] = DirichletResidual( domain=pde.boundaries[boundary], model_type="x", f_rhs=dirichlet_bc_at_t, ) return pde params = { "Implicit Euler (1st order)": { "nt": 100, "tableau": build_implicit_euler_tableau(), "explicit_pde": None, "implicit_pde": pde_factory, }, "Alexander (3rd order)": { "nt": 50, "tableau": build_alexander_tableau(), "explicit_pde": None, "implicit_pde": pde_factory, }, } key = jax.random.PRNGKey(0) space = None spaces = [] losses = [] for method, param in params.items(): discrete_pinn = DiscretePINN( DOM_X, DOM_T, SAMPLER, 1, param["nt"], param["tableau"], param["explicit_pde"], param["implicit_pde"], exact_solution=exact_sol, ) if space is None: # train the initial condition only once, and share it across methods nn = MLP(in_size=1, out_size=1, hidden_sizes=[12] * 3, key=key) space = ApproximationSpace({"x": 1}, [(nn, "scalar", None)], model_type="x") print("Initializing the discrete PINN...") key, discrete_pinn = discrete_pinn.initialize( key, space, f_init, N_EPOCHS_INIT, N_COLLOC, file_name="discrete_pinn_1d_heat_variable_diffusion_init", retrain=False, linesearch="strong-wolfe", adaptive_matrix_regularization=True, ) space = discrete_pinn.space print("Initializing the discrete PINN... Done\n") plot_abstract_approx_spaces( [space], DOM_X, solution=f_init, error=f_init, title="initial condition" ) print(f"\nSolving with the {method} method...") key, discrete_pinn = discrete_pinn.solve( key, space, N_EPOCHS, N_COLLOC, n_bc_colloc=N_BC_COLLOC ) print(f"Solving with the {method} method... Done\n") key, l2, linf = discrete_pinn.compute_relative_error( key, discrete_pinn.space, DOM_T[1] ) print(f"{method}: relative L2 error = {l2:.2e}, relative Linf error = {linf:.2e}") spaces.append(discrete_pinn.space) losses.append(discrete_pinn.losses) plot_abstract_approx_spaces( spaces, DOM_X, loss=losses, solution=lambda x: exact_sol(jnp.ones_like(x) * DOM_T[1], x), error=lambda x: exact_sol(jnp.ones_like(x) * DOM_T[1], x), title=f"solution at final time t={DOM_T[1]}", titles=list(params.keys()), ) plt.show()