r"""Solves the heat equation in 1D with time-discrete kernel collocation. .. math:: \partial_t u - \partial_{xx} u & = 0 \text{ in } (0, 1) \times (0, T) \\ u & = 0 \text{ on } \{0, 1\} \times (0, T) \\ u & = \sin(\pi x) \text{ on } (0, 1) \times \{0\} whose exact solution is :math:`u(x, t) = e^{-\pi^2 t} \sin(\pi x)`. The solution is expanded on a Gaussian-kernel basis centered on ``N_POINTS`` equispaced points, which are also the collocation points; the kernel width is tied to the point spacing (``sigma = SIGMA_FACTOR * spacing``). A narrow kernel (``SIGMA_FACTOR`` close to 1, fine for a reaction term, see ``reaction_1d_matrix_free.py``) makes the second-derivative Jacobian of the Laplacian badly conditioned: ``SIGMA_FACTOR`` around 2-3 is needed here, and not every ``N_POINTS`` is equally well conditioned for a given factor. The time integration uses the Pareschi-Russo Runge-Kutta scheme (second order, stiffly accurate) through ``TimeDiscreteCollocationScheme``, with ``NT`` time steps. Dirichlet is imposed weakly, through the boundary collocation residual: collocation has no notion of a "boundary DOF" to lift, only boundary collocation points. Each stage is a dense Newton solve: collocating the Laplacian with a Gaussian-RBF basis gives a Jacobian whose condition number grows with the point count, which erases the speed advantage of the matrix-free path (see ``reaction_1d_matrix_free.py``). The time and space convergence of the three schemes (explicit/implicit Euler, Pareschi-Russo) on this problem are studied in ``comparative_studies_jax/kernel/heat_1d_segment_convergence.py``, and measured by the ``heat_1d`` datasets of ``benchmarks/benchmarks_jax/kernel``. """ 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.linear_approximation.basis.kernel_basis import KernelBasis from scimba_jax.linear_approximation.basis.kernel_function import GaussianKernel from scimba_jax.linear_approximation.collocation.time_discrete_collocation_scheme import ( TimeDiscreteCollocationScheme, ) from scimba_jax.linear_approximation.variables.collocation_variables import ( CollocationVariables, ) from scimba_jax.physical_models.elliptic_pde.laplacians import LaplacianResidual from scimba_jax.time_discrete.butcher_tableau import build_pareschi_russo_tableau T_FINAL = 0.1 N_POINTS = 20 SIGMA_FACTOR = 3.0 NT = 8 DT = T_FINAL / NT N_EVAL = 400 def u0(x, mu=None): """Initial condition u(x, 0) = sin(pi x).""" return jnp.sin(jnp.pi * x) def u_exact(t, x): """Exact solution u(x, t) = exp(-pi^2 t) sin(pi x).""" return jnp.exp(-(jnp.pi**2) * t) * jnp.sin(jnp.pi * x) # ── Kernel collocation space ───────────────────────────────────────────────── domain = Segment1D(jnp.array([[0.0, 1.0]]), is_main_domain=True) x_centers = jnp.linspace(0.0, 1.0, N_POINTS).reshape(-1, 1) x_centers_bc = jnp.array([[0.0], [1.0]]) sigma = SIGMA_FACTOR / (N_POINTS - 1) basis = KernelBasis( dim=1, output_dim=1, kernel_function=GaussianKernel(sigma=sigma), centers=x_centers, basis_type="scalar", ) variables = CollocationVariables(basis=basis, nb_variables=1) # ── Time-discrete scheme ───────────────────────────────────────────────────── # Time-independent data (homogeneous heat equation, zero boundary): the same # LaplacianResidual is reused, unchanged, at every stage and every time step. scheme = TimeDiscreteCollocationScheme( spatial_residual_factory=lambda t: LaplacianResidual( domain=domain, f_rhs=lambda x, mu: jnp.zeros(()) ), variables=variables, collocation_points=x_centers, bc_collocation_points=x_centers_bc, main_domain=domain, butcher_tableau=build_pareschi_russo_tableau(), dt=DT, dirichlet_factory=lambda t: (lambda x, n, mu: jnp.zeros(1)), ) print( f"Solving the 1D heat equation with kernel collocation: n_points={N_POINTS}, " f"sigma_factor={SIGMA_FACTOR}, nt={NT}, dt={DT:.2e} -- Pareschi-Russo" ) start = timeit.default_timer() dofsl_init = scheme.initialize(u0) dofsl_final, history = scheme.solve(dofsl_init, t0=0.0, nt=NT) # The solve is dispatched asynchronously: block before reading the clock. jax.block_until_ready(history) end = timeit.default_timer() print(f"time for {NT} time steps: {end - start:.2f}s") # ── Error against the exact solution ───────────────────────────────────────── xs = jnp.linspace(0.0, 1.0, N_EVAL) def u_h_at(dofsl): """u_h(xs) for the given DOFs.""" variables.dofsl = dofsl return jax.vmap(lambda x: variables.evaluate(x[None]))(xs).reshape(-1) u_h = u_h_at(dofsl_final) u_ex = u_exact(T_FINAL, xs) # Collocation has no mesh/quadrature machinery: the L2 norm is estimated by # trapezoidal quadrature on a dense evaluation grid. rel_l2 = jnp.sqrt(jnp.trapezoid((u_h - u_ex) ** 2, xs) / jnp.trapezoid(u_ex**2, xs)) print(f"rel. L2 error at t={T_FINAL}: {rel_l2:.3e}") print(f"max abs error at t={T_FINAL}: {jnp.max(jnp.abs(u_h - u_ex)):.3e}") # ── Plot: solution and pointwise error at a few time steps ─────────────────── fig, (ax_u, ax_err) = plt.subplots(1, 2, figsize=(11, 4.5)) for step in (0, NT // 2, NT): t = step * DT u_h_t = u_h_at(history[step]) u_ex_t = u_exact(t, xs) (line,) = ax_u.plot(xs, u_h_t, "--", label=f"collocation, t={t:.3f}") ax_u.plot( xs, u_ex_t, "-", color=line.get_color(), alpha=0.4, linewidth=3, label=f"exact, t={t:.3f}", ) ax_err.semilogy(xs, jnp.abs(u_h_t - u_ex_t), label=f"t={t:.3f}") ax_u.set_xlabel("x") ax_u.set_ylabel("u") ax_u.set_title("solution") handles, labels = ax_u.get_legend_handles_labels() # Two columns, collocation then exact, to keep the legend below the curves. order = list(range(0, len(handles), 2)) + list(range(1, len(handles), 2)) ax_u.legend( [handles[i] for i in order], [labels[i] for i in order], loc="lower center", ncol=2, fontsize="small", ) ax_err.set_xlabel("x") ax_err.set_ylabel(r"$|u_h - u_\mathrm{exact}|$") ax_err.set_title(f"pointwise error (rel. L2 at t={T_FINAL}: {rel_l2:.2e})") ax_err.legend() fig.suptitle( "1D heat equation, kernel collocation " f"(Pareschi-Russo, {N_POINTS} points, {NT} time steps)" ) fig.tight_layout() plt.show()