r"""1D parametric heat equation with TimeDiscreteFEscheme, 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). FEM counterpart of ``examples_jax/dg/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 ``TimeDiscreteFEscheme``/ ``EllipticFEscheme`` instead of the DG scheme. ``DiffusionWeakForm`` is simpler here than its DG twin: FEM imposes Dirichlet by lifting DOFs (not through a numerical flux), so no ``get_fields``/``fields["A"]`` override is needed -- ``bilinear_form`` alone (``alpha * grad(u).grad(v)``, read live off ``self.alpha``, a genuine pytree child -- see the DG file's module docstring for why that matters under vmap) is enough. """ import functools 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.galerkin.fem.time_discrete_fe_scheme import ( TimeDiscreteFEscheme, ) from scimba_jax.linear_approximation.meshes.cartesian_mesh import cartesian_mesh from scimba_jax.linear_approximation.variables.variables_fe import VariablesFE from scimba_jax.physical_models.abstract_weak_form import AbstractWeakForm from scimba_jax.time_discrete.butcher_tableau import build_implicit_euler_tableau DIM = 1 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_CELLS = 20 ORDER = 2 NT = 100 DT = T_FINAL / NT N_PLOT = 200 X_PLOT = jnp.linspace(X_MIN, X_MAX, N_PLOT)[:, None] def make_variables(n_cells, order=ORDER, quad_order=4): mesh = cartesian_mesh(n_cells=[n_cells], quad_order=quad_order) basis = AnalyticBasis( nb_basis=order + 1, out_dim=1, mesh=mesh, local_basis=lambda y, i, m, k=order: local_lagrange_basis( y, i, m, order=k, out_dim=1 ), basis_type="scalar", ) return VariablesFE(basis=basis, nb_variables=1) def u0(x): """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 DiffusionWeakForm(AbstractWeakForm): r"""Weak form for ``du/dt - 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 ``bilinear_form`` -- not baked into a closure at ``__init__`` time -- so it tracks whichever ``alpha`` a given ``jax.vmap`` slot carries. See the DG twin of this file for why that distinction matters (there, the same ``alpha`` also has to reach ``SIPGFlux`` through ``fields["A"]``; here FEM has no interface flux, so ``bilinear_form``/``linear_form`` are enough). """ def __init__(self, dim: int, alpha: jnp.ndarray): super().__init__(dim=dim) self.alpha = jnp.asarray(alpha) def bilinear_form(self, u, v): return self.alpha * u.gradient("x").dot(v.gradient("x")) def linear_form(self, v): return v * 0.0 def _solve_one_draw( n_cells: int, order: int, dt: float, nt: int, x_plot: jnp.ndarray, alpha: jnp.ndarray, ) -> tuple[jnp.ndarray, jnp.ndarray, jnp.ndarray, jnp.ndarray]: """One FE 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 = make_variables(n_cells, order) weak_form = DiffusionWeakForm(dim=DIM, alpha=alpha) scheme = TimeDiscreteFEscheme( spatial_weak_form_factory=weak_form, variables=variables, butcher_tableau=build_implicit_euler_tableau(), dt=dt, dirichlet=lambda x: 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 = variables.evaluate(x_plot)[:, 0] 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.mean(abs_err**2)) / jnp.sqrt(jnp.mean(u_ex**2)) return u_h, u_ex, abs_err, rel_l2_err def run_batch( n_cells: int, order: int, 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_cells, order, dt, nt, x_plot) return jax.vmap(solve_one)(alpha_batch) def compare_batching_time( n_cells: int, order: int, 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 THIS file's comparison (2000 draws timing FASTER than 1: FE's per-draw compute is cheap enough, at this problem size, for the warm-up to dominate and invert the expected ordering). 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(2, 1, 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_cells, order, 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_cells, order, 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"FE, n_cells={N_CELLS}, order={ORDER}, 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 FE 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="FE 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( f"FE, parametric heat equation, alpha ~ Uniform([{ALPHA_MIN}, {ALPHA_MAX}]), " f"{N_RUNS_PARAMETRIC} draws, n_cells={N_CELLS}, order={ORDER}" ) 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_CELLS, ORDER, 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)