r"""1D parametric heat equation with TimeDiscreteCollocationScheme, vmapped over alpha. du/dt - alpha * u_xx = 0 on [0, 1], u = 0 on the boundary u(x, 0) = sin(pi x) Exact solution: u(x, t; alpha) = exp(-alpha * pi^2 * t) sin(pi x). Collocation counterpart of ``examples_jax/dg/time_dependent/1d_heat_equation_parametric.py`` and ``examples_jax/fem/solve/time_dependent/ 1d_heat_equation_parametric.py`` -- same physics, same batching pattern (a physical parameter ``alpha``, uniformly drawn ``N_RUNS_PARAMETRIC`` times and solved in one ``jax.vmap`` call, following the ``parametric`` rows of ``benchmarks/benchmarks_jax/dg_enriched_well_balanced/``), through ``TimeDiscreteCollocationScheme``/ ``EllipticCollocationScheme`` (dense Newton, Gaussian-RBF kernel basis) instead of a mesh-based scheme. ``DiffusionResidual`` mirrors ``ShallowWaterWeakForm``'s ``z_max``/``sigma`` pattern: ``alpha`` is a plain ``jnp.ndarray`` attribute (a genuine, live pytree child), read directly in ``construct_residual`` rather than captured by a closure built once in ``__init__`` -- see the DG file's module docstring for why the latter would go stale across the vmap batch. """ import functools import time import jax import jax.numpy as jnp import matplotlib.pyplot as plt import numpy as np 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.abstract_residuals import InteriorResidual from scimba_jax.time_discrete.butcher_tableau import build_implicit_euler_tableau X_MIN, X_MAX = 0.0, 1.0 T_FINAL = 0.1 ALPHA_MIN, ALPHA_MAX = 1.0, 1.5 N_RUNS_PARAMETRIC = 2000 N_POINTS = 20 SIGMA_FACTOR = 3.0 NT = 100 DT = T_FINAL / NT N_PLOT = 200 X_PLOT = jnp.linspace(X_MIN, X_MAX, N_PLOT)[:, None] def make_variables(n_points, sigma_factor=SIGMA_FACTOR): """Build a Gaussian-kernel collocation space on ``n_points`` centers over ``[0, 1]``, see ``1d_heat_equation.py`` (this module's non-parametric twin) for why ``sigma_factor`` needs to sit around 2-3 for a well-conditioned dense Newton solve of the Laplacian. """ spacing = 1.0 / (n_points - 1) sigma = sigma_factor * spacing x_centers = jnp.linspace(0.0, 1.0, n_points).reshape(-1, 1) x_centers_bc = jnp.array([[0.0], [1.0]]) domain = Segment1D(jnp.array([[0.0, 1.0]]), is_main_domain=True) 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) return variables, domain, x_centers, x_centers_bc def u0(x, mu=None): """Initial condition u(x, 0) = sin(pi x), alpha-independent.""" return jnp.sin(jnp.pi * x) def u_exact(alpha, t, x): """Exact solution u(x, t; alpha) = exp(-alpha * pi^2 * t) sin(pi x).""" return jnp.exp(-alpha * jnp.pi**2 * t) * jnp.sin(jnp.pi * x) class DiffusionResidual(InteriorResidual): r"""Residual of ``-alpha * lap(u) = 0``. ``self.alpha`` is a plain ``jnp.ndarray`` attribute, auto-classified a dynamic pytree child (see ``ScimbaPytree``/``_is_dynamic_value``), read live in ``construct_residual`` -- exactly ``ShallowWaterWeakForm``'s ``z_max``/``sigma`` pattern -- so it tracks whichever ``alpha`` a given ``jax.vmap`` slot carries. """ def __init__(self, domain, alpha: jnp.ndarray, f_rhs=None): super().__init__(domain=domain, size=1, model_type="x", f_rhs=f_rhs) self.alpha = jnp.asarray(alpha) def construct_residual(self, *variables): u = variables[0] return -self.alpha * u.laplacian("x") def _solve_one_draw( n_points: int, sigma_factor: float, dt: float, nt: int, x_plot: jnp.ndarray, alpha: jnp.ndarray, ) -> tuple[jnp.ndarray, jnp.ndarray, jnp.ndarray, jnp.ndarray]: """One collocation solve at diffusivity ``alpha``, kept fully traceable (no ``float(...)``/``np.asarray(...)``) so :func:`run_batch` can ``jax.vmap`` it over a whole batch of draws -- see ``shallow_water.py``'s ``_solve_one_draw`` for the same contract. Returns: ``(u_h, u_ex, abs_err, rel_l2_err)`` at ``t = T_FINAL``, evaluated on ``x_plot`` (``u_h``/``u_ex``/``abs_err`` shape ``(n_plot,)``), plus the scalar relative L2 error. """ variables, domain, x_centers, x_centers_bc = make_variables(n_points, sigma_factor) def spatial_residual_factory(t): return DiffusionResidual( domain=domain, alpha=alpha, f_rhs=lambda x, mu: jnp.zeros(()) ) scheme = TimeDiscreteCollocationScheme( spatial_residual_factory=spatial_residual_factory, variables=variables, collocation_points=x_centers, bc_collocation_points=x_centers_bc, main_domain=domain, butcher_tableau=build_implicit_euler_tableau(), dt=dt, dirichlet_factory=lambda t: (lambda x, n, mu: jnp.zeros(1)), ) dofsl_init = scheme.initialize(u0) dofsl_final, _ = scheme.solve(dofsl_init, t0=0.0, nt=nt) variables.dofsl = dofsl_final u_h = jax.vmap(lambda x: variables.evaluate(x[None]))(x_plot[:, 0]).reshape(-1) u_ex = jax.vmap(lambda x: u_exact(alpha, T_FINAL, x))(x_plot[:, 0]) abs_err = jnp.abs(u_h - u_ex) rel_l2_err = jnp.sqrt(jnp.trapezoid(abs_err**2, x_plot[:, 0])) / jnp.sqrt( jnp.trapezoid(u_ex**2, x_plot[:, 0]) ) return u_h, u_ex, abs_err, rel_l2_err def run_batch( n_points: int, sigma_factor: float, dt: float, nt: int, x_plot: jnp.ndarray, alpha_batch: jnp.ndarray, ) -> tuple[jnp.ndarray, jnp.ndarray, jnp.ndarray, jnp.ndarray]: """Solves every draw of ``alpha_batch`` in a SINGLE compiled program via ``jax.vmap`` -- see the module docstring and ``shallow_water.py``'s ``run_category_batch`` for why this beats a Python loop of separate :func:`_solve_one_draw` calls. """ solve_one = functools.partial( _solve_one_draw, n_points, sigma_factor, dt, nt, x_plot ) return jax.vmap(solve_one)(alpha_batch) def compare_batching_time( n_points: int, sigma_factor: float, dt: float, nt: int, x_plot: jnp.ndarray, alpha_batch: jnp.ndarray, ) -> tuple[jnp.ndarray, jnp.ndarray, jnp.ndarray, jnp.ndarray]: """Times 1 draw against the full ``N_RUNS_PARAMETRIC``-draw vmapped batch, to make explicit the point ``shallow_water.py``'s own parametric study makes: ``jax.vmap`` pays the one-time jit compile ONCE, shared by every draw, so per-draw cost drops sharply as the batch grows -- a lone solve pays that same compile for a single result. Both calls are freshly traced (a batch of 1 draw is a different input shape than a batch of ``N_RUNS_PARAMETRIC``, hence its own compile), so each elapsed time genuinely includes its own compile-plus-run cost -- but that alone is not enough for a fair comparison: the FIRST call in a process also silently pays a one-time warm-up (XLA backend init, first construction of scimba's pytree-registered classes, ...) worth a second or more, unrelated to this problem's size. Left uncontrolled, that makes whichever measurement runs first look artificially slow -- observed to flip the FE version of this file's comparison (2000 draws timing FASTER than 1). An untimed throwaway solve, on a shape neither measurement below reuses, absorbs that one-time cost up front. Returns: The full batch's ``(u_h, u_ex, abs_err, rel_l2_err)``, so the caller can reuse it instead of solving the full batch a second time. """ n_runs = alpha_batch.shape[0] print("warmup ...") warmup_out = run_batch(4, sigma_factor, T_FINAL / 2, 2, x_plot, alpha_batch[:1]) jax.block_until_ready(warmup_out) print("timing 1 run vs. the full vmapped batch ...") t0 = time.perf_counter() one_out = run_batch(n_points, sigma_factor, dt, nt, x_plot, alpha_batch[:1]) jax.block_until_ready(one_out) elapsed_one = time.perf_counter() - t0 print( f" 1 run : {elapsed_one:.3f}s ({1e3 * elapsed_one:.2f} ms/draw)" ) t0 = time.perf_counter() batch_out = run_batch(n_points, sigma_factor, dt, nt, x_plot, alpha_batch) jax.block_until_ready(batch_out) elapsed_batch = time.perf_counter() - t0 per_draw_batch = elapsed_batch / n_runs print( f" {n_runs} runs (vmapped) : {elapsed_batch:.3f}s " f"({1e3 * per_draw_batch:.3f} ms/draw)" ) print( f" amortized speedup per draw : {elapsed_one / per_draw_batch:.1f}x " "(one jit compile shared by every draw in the batch)" ) return batch_out def print_summary(alpha_batch: jnp.ndarray, rel_l2_err_batch: jnp.ndarray) -> None: alpha = np.asarray(alpha_batch) err = np.asarray(rel_l2_err_batch) print( f"\n{len(err)} draws of alpha ~ Uniform([{ALPHA_MIN}, {ALPHA_MAX}]), " f"collocation, n_points={N_POINTS}, sigma_factor={SIGMA_FACTOR}, nt={NT}" ) print(f" alpha mean={alpha.mean():.3f} std={alpha.std():.3f}") print( f" rel. L2 error at t={T_FINAL} mean={err.mean():.3e} std={err.std():.3e}" f" min={err.min():.3e} max={err.max():.3e}" ) def plot_results( x_plot: jnp.ndarray, u_h_batch: jnp.ndarray, u_ex_batch: jnp.ndarray, abs_err_batch: jnp.ndarray, ) -> None: """Mean +/- std, over the batch of alpha draws, of the collocation solution and of the pointwise error against the exact solution, both at ``t = T_FINAL``. """ x = np.asarray(x_plot[:, 0]) u_h_batch = np.asarray(u_h_batch) u_ex_batch = np.asarray(u_ex_batch) abs_err_batch = np.asarray(abs_err_batch) mean_u_h, std_u_h = u_h_batch.mean(axis=0), u_h_batch.std(axis=0) mean_u_ex, std_u_ex = u_ex_batch.mean(axis=0), u_ex_batch.std(axis=0) mean_err, std_err = abs_err_batch.mean(axis=0), abs_err_batch.std(axis=0) fig, axes = plt.subplots(1, 2, figsize=(12, 5)) axes[0].plot(x, mean_u_ex, "k--", linewidth=1.5, label="exact, mean over alpha") axes[0].fill_between( x, mean_u_ex - std_u_ex, mean_u_ex + std_u_ex, color="k", alpha=0.15 ) axes[0].plot(x, mean_u_h, label="collocation scheme, mean over alpha") axes[0].fill_between( x, mean_u_h - std_u_h, mean_u_h + std_u_h, alpha=0.25, label="+/- 1 std" ) axes[0].set_xlabel("x") axes[0].set_ylabel(f"u(x, t={T_FINAL})") axes[0].set_title("Solution: mean +/- std over alpha draws") axes[0].legend(fontsize=8) axes[0].grid(True, alpha=0.3) axes[1].plot(x, mean_err, label="|u_h - u_exact|, mean over alpha") axes[1].fill_between( x, np.clip(mean_err - std_err, 1e-16, None), mean_err + std_err, alpha=0.25, label="+/- 1 std", ) axes[1].set_xlabel("x") axes[1].set_ylabel("pointwise error") axes[1].set_yscale("log") axes[1].set_title("Pointwise error: mean +/- std over alpha draws") axes[1].legend(fontsize=8) axes[1].grid(True, alpha=0.3, which="both") fig.suptitle( "Collocation, parametric heat equation, " f"alpha ~ Uniform([{ALPHA_MIN}, {ALPHA_MAX}]), " f"{N_RUNS_PARAMETRIC} draws, n_points={N_POINTS}" ) fig.tight_layout() plt.show() if __name__ == "__main__": key = jax.random.PRNGKey(0) alpha_batch = jax.random.uniform( key, (N_RUNS_PARAMETRIC,), minval=ALPHA_MIN, maxval=ALPHA_MAX ) u_h_batch, u_ex_batch, abs_err_batch, rel_l2_err_batch = compare_batching_time( N_POINTS, SIGMA_FACTOR, DT, NT, X_PLOT, alpha_batch ) print_summary(alpha_batch, rel_l2_err_batch) plot_results(X_PLOT, u_h_batch, u_ex_batch, abs_err_batch)