"""The 1-D heat equation solved in SPACE-TIME, on a tensorised mesh. d_t u - nu d_xx u = 0 on (0, 1) x (0, T) u(x, 0) = u0(x) the initial condition u(0, t) = u(1, t) = 0 the boundary condition Time is not stepped. The mesh is a 1-D mesh in ``x`` times a 1-D mesh in ``t``, so the whole space-time slab is ONE 2-D domain, and the problem is one elliptic solve over it. What that buys is the point of the example: **the initial condition stops being a special object.** It is a Dirichlet condition on the ``t = 0`` side, written and enforced exactly like the two spatial ones -- three sides of the square carry data, the fourth carries nothing. **The weak form is a space-time one, and it is NOT the 2-D Laplacian.** Only the SPACE derivative is integrated by parts:: a(u, v) = int_slab (d_t u) v + nu int_slab (d_x u)(d_x v) Two things follow, both of which matter: * it is **not symmetric** -- the time term is a first derivative, so the operator is a degenerate advection-diffusion (``A = diag(nu, 0)``, ``b = e_t``) and CG does not apply to it. The assembled path is used, not matrix-free CG; * leaving ``d_t`` unintegrated is exactly what makes ``t = T`` need NO condition. Integrating it by parts would produce a boundary term there and the slab would ask for data at the final time, which one does not have. The asymmetry of the formulation IS the arrow of time. **The control** is the exact solution, by Fourier: with ``u0`` expanded on ``sin(k pi x)``, each mode decays as ``exp(-nu k^2 pi^2 t)``, so the reference owes nothing to the code. The Gaussian initial datum is centred and narrow enough that its value at the walls, ``exp(-0.5 * (0.5 / 0.07)^2) = 8e-12``, is far below the discretisation error -- otherwise the corners ``(0, 0)`` and ``(1, 0)`` would carry a genuine incompatibility between the initial and the boundary data, and the solution would have a corner singularity there rather than a small one. """ # %% import time import jax import jax.numpy as jnp import matplotlib.pyplot as plt import numpy as np from scimba_jax.linear_approximation.basis.analytic_bases import local_lagrange_basis from scimba_jax.linear_approximation.basis.general_bases import AnalyticBasis from scimba_jax.linear_approximation.error_analysis import l2_error from scimba_jax.linear_approximation.galerkin.fem.elliptic_fe_scheme import ( EllipticFEscheme, ) from scimba_jax.linear_approximation.meshes.mesh import Mesh from scimba_jax.linear_approximation.meshes.tensor_mesh import tensor_mesh from scimba_jax.linear_approximation.quad.gauss_quad import UnitSquareTensorized from scimba_jax.linear_approximation.variables.variables_fe import VariablesFE from scimba_jax.mapping.mapping import InvertibleFunction, Mapping from scimba_jax.physical_models.abstract_physical_weak_model import ( AbstractPhysicalWeakModel, ) from scimba_jax.physical_models.classical_weakform.diffusion_advection_reaction_weak_form import ( # noqa: E501 EllipticWeakForm, ) from scimba_jax.physical_models.weak_boundary_conditions import Dirichlet DIM = 2 POLY_ORDER = 2 QUAD_ORDER = 2 * POLY_ORDER + 1 OUT_DIM = 1 N_SPACE = 32 N_TIME = 12 # ⚠ A SHORT slab, which is how space-time formulations are actually used: one # solves a thin slice and marches, rather than putting the whole history into # one system. # # What decides whether anything is VISIBLE over so short a time is not `T` but # `nu T / sigma^2`, which is 1.0 here: a tight enough initial bump diffuses # appreciably in very little time. It ends `sqrt(1 + 2 nu T / sigma^2) = 1.7` # times wider and as much lower. Widening `T` instead would have bought the # same picture at several times the cost. # # ⚠ And tightening `sigma` has its own floor: at 32 cells and Q2 there are 2.2 # cells per `sigma`, which resolves it. Below that the error would be measuring # the sampling rather than the scheme -- the trap already paid for once on the # tokamak blob. FINAL_TIME = 0.1 NU = 0.05 BLOB_CENTRE = 0.5 BLOB_WIDTH = 0.07 N_MODES = 200 # ── The physics ───────────────────────────────────────────────────────────── def initial_condition(x): """A Gaussian bump, well inside the walls.""" return jnp.exp(-0.5 * ((x - BLOB_CENTRE) / BLOB_WIDTH) ** 2) class SpaceTimeHeat(EllipticWeakForm): """``a(u, v) = int (d_t u) v + nu int (d_x u)(d_x v)`` on the slab. ⚠ The two derivatives are treated DIFFERENTLY, and that is the physics, not a shortcut. Space is integrated by parts -- symmetric, and its boundary term is what carries the wall condition. Time is not -- so no term appears at ``t = T``, which is why the final time needs no data and why the operator is not symmetric. Written as a degenerate advection-diffusion because that is what a space-time slab IS: the diffusion tensor is ``diag(nu, 0)``, blind to the time direction, and the advection is the unit vector ALONG it. The parent's generic form would compute exactly this; ``bilinear_form`` is overridden only to spare the traced graph a 2x2 matmul against a matrix that is mostly zero, the same reason ``LaplacianWeakForm`` overrides it. """ def __init__(self, dim: int, nu: float): super().__init__( dim=dim, A=lambda x: jnp.diag(jnp.array([nu, 0.0])), b=lambda x: jnp.array([0.0, 1.0]), c=lambda x: jnp.zeros(()), f=lambda x: jnp.zeros(OUT_DIM), ) self.nu = nu def bilinear_form(self, u, v): grad_u, grad_v = u.gradient("x"), v.gradient("x") transport = grad_u.component(1) * v diffusion = (grad_u.component(0) * grad_v.component(0)) * self.nu return transport + diffusion class HeatSlab(AbstractPhysicalWeakModel): """The slab and its three sides of data; the fourth is the final time. Keyed by side name, exactly as a PINN model keys its residuals by boundary label -- and the initial condition is just one of the entries. ``north`` is absent on purpose: at ``t = T`` the space-time form asks for nothing. """ def __init__(self, nu=NU): super().__init__(dim=DIM) self.weak_forms = {"interior": SpaceTimeHeat(dim=DIM, nu=nu)} self.boundary_conditions = { # t = 0: the initial condition, as a Dirichlet condition. "south": Dirichlet(lambda x: jnp.array([initial_condition(x[0])])), # x = 0 and x = 1: the walls. "west": Dirichlet(lambda x: jnp.zeros(OUT_DIM)), "east": Dirichlet(lambda x: jnp.zeros(OUT_DIM)), } # ── The reference, by Fourier ─────────────────────────────────────────────── def _mode_amplitudes(n_modes=N_MODES, n_quad=4000): """``a_k = 2 int_0^1 u0(x) sin(k pi x) dx``, by a fine trapezium rule. Fine enough to be exact to rounding for a smooth datum: with 4 000 points the rule's error is far below the discretisation being measured, so this is a REFERENCE and not a second approximation to compare against. """ x = jnp.linspace(0.0, 1.0, n_quad) modes = jnp.arange(1, n_modes + 1) integrand = initial_condition(x)[None, :] * jnp.sin( modes[:, None] * jnp.pi * x[None, :] ) return 2.0 * jnp.trapezoid(integrand, x, axis=1) _AMPLITUDES = _mode_amplitudes() def u_exact(point): """``sum_k a_k exp(-nu k^2 pi^2 t) sin(k pi x)``, at one space-time point.""" x, t = point[0], point[1] modes = jnp.arange(1, _AMPLITUDES.shape[0] + 1) decay = jnp.exp(-NU * (modes * jnp.pi) ** 2 * t) return jnp.array([jnp.sum(_AMPLITUDES * decay * jnp.sin(modes * jnp.pi * x))]) # ── The mesh: 1-D in x, times 1-D in t ────────────────────────────────────── def _line(n_cells, length): """A 1-D mesh of ``[0, length]``.""" return Mesh( dim=1, n_cells=(n_cells,), ref_quad=UnitSquareTensorized(dim=1, order=QUAD_ORDER), mapping=Mapping( mappings=[InvertibleFunction(lambda x: length * x, lambda y: y / length)] ), ) def slab(n_space=N_SPACE, n_time=N_TIME): """``[0,1]_x`` times ``[0,T]_t``, named as the PINNs would name it.""" return tensor_mesh(_line(n_space, 1.0), _line(n_time, FINAL_TIME), names=("x", "t")) def make_scheme(n_space=N_SPACE, n_time=N_TIME, nu=NU): mesh = slab(n_space, n_time) basis = AnalyticBasis( nb_basis=(POLY_ORDER + 1) ** DIM, out_dim=OUT_DIM, mesh=mesh, local_basis=lambda c, i, m: local_lagrange_basis( c, i, m, order=POLY_ORDER, out_dim=OUT_DIM ), basis_type="scalar", ) return EllipticFEscheme( HeatSlab(nu), VariablesFE(basis=basis, nb_variables=OUT_DIM) ) # ── Reading the slab back ─────────────────────────────────────────────────── def sample(scheme, n_x=120, n_t=90): """The solution on a regular space-time grid, for the picture.""" xs = np.linspace(0.0, 1.0, n_x) ts = np.linspace(0.0, FINAL_TIME, n_t) grid = np.stack(np.meshgrid(xs, ts, indexing="ij"), -1).reshape(-1, 2) values = jax.vmap(lambda p: scheme.variables.local_evaluate(p)[0])( jnp.asarray(grid) ) return xs, ts, np.asarray(values).reshape(n_x, n_t) def main(): print( f"FEM espace-temps Q{POLY_ORDER}, maillage {N_SPACE} x {N_TIME} " f"= {N_SPACE * N_TIME} mailles du slab (0,1) x (0,{FINAL_TIME})" ) print(f" bords du maillage : {sorted(slab().boundary_groups)}") started = time.perf_counter() scheme = EllipticFEscheme.solve(make_scheme(), matrix_free=False, max_iter=1) jax.block_until_ready(scheme.variables.dofsl) print(f" solve {time.perf_counter() - started:6.2f} s") print( f" erreur L2 relative sur le slab : {float(l2_error(scheme, u_exact, relative=True)):.3e}" ) # ⚠ The rate is measured in the SLAB norm, so it mixes the two directions. # Both are refined together here; refining one alone is what would separate # them, and `test_tensor_mesh.py` already does that for the projection. print("\n convergence (les deux directions raffinees ensemble) :") previous = None for factor in (1, 2): refined = EllipticFEscheme.solve( make_scheme(N_SPACE * factor, N_TIME * factor), matrix_free=False, max_iter=1, ) error = float(l2_error(refined, u_exact, relative=True)) rate = "" if previous is None else f" taux {np.log2(previous / error):.2f}" print(f" {N_SPACE * factor:3d} x {N_TIME * factor:3d} : {error:.3e}{rate}") previous = error xs, ts, computed = sample(scheme) exact = np.asarray( jax.vmap(u_exact)( jnp.asarray(np.stack(np.meshgrid(xs, ts, indexing="ij"), -1).reshape(-1, 2)) ) ).reshape(len(xs), len(ts)) figure, axes = plt.subplots(1, 3, figsize=(15, 4.2)) levels = dict(vmin=0.0, vmax=float(max(computed.max(), exact.max())), cmap="turbo") for panel, (values, title) in enumerate( ((computed, "FEM espace-temps"), (exact, "exact (Fourier)")) ): image = axes[panel].pcolormesh(ts, xs, values, shading="gouraud", **levels) figure.colorbar(image, ax=axes[panel], fraction=0.046) axes[panel].set_title(title) axes[panel].set_xlabel("t") axes[panel].set_ylabel("x") difference = axes[2].pcolormesh( ts, xs, computed - exact, shading="gouraud", cmap="coolwarm" ) figure.colorbar(difference, ax=axes[2], fraction=0.046) axes[2].set_title(f"ecart (max {np.abs(computed - exact).max():.1e})") axes[2].set_xlabel("t") axes[2].set_ylabel("x") figure.suptitle( "chaleur 1D en formulation espace-temps : maillage 1D en x tensorise " "avec un maillage 1D en t" ) plt.tight_layout() plt.show() if __name__ == "__main__": main()