r"""1D parametric heat equation with TimeDiscreteDGscheme, 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). Parametric counterpart of the DG rows of ``benchmarks/benchmarks_jax/ heat_time_convergence`` (same problem, same SIPG setup): instead of a single fixed diffusivity, ``alpha`` is a *physical parameter* drawn uniformly in ``[ALPHA_MIN, ALPHA_MAX]``, and ``N_RUNS_PARAMETRIC`` independent draws are solved in a SINGLE compiled program via ``jax.vmap`` -- exactly the pattern the ``parametric`` rows of ``benchmarks/benchmarks_jax/dg_enriched_well_balanced/`` use to batch many bump shapes ``(z_max, sigma)`` (see its ``benchmark_utils/well_balanced.py`` for why one traced-once ``vmap`` call beats a Python loop of separate solves, even a non-recompiling one). ``alpha`` must be a genuine pytree CHILD of the weak form (an array attribute, not a value baked into a closure at ``__init__`` time and then read back through the weak form's own auto-classified static ``aux_data`` -- see CLAUDE.md's note on ``ShallowWaterWeakForm``'s ``z_max``/``sigma`` for the exact same requirement). ``DiffusionWeakForm`` below follows that pattern: ``self.alpha`` is a plain ``jnp.ndarray`` attribute (auto-classified dynamic, i.e. a pytree child), read directly -- not through a stored closure -- both in ``bilinear_form`` (the volume term, following ``LaplacianWeakForm``'s own override to skip the ``A=I`` matmul) and in a ``get_fields`` override that hands ``SIPGFlux`` a live ``alpha * I`` diffusion matrix (``fields["A"]``) for the face terms -- reusing ``EllipticWeakForm``'s static-``A``-closure machinery instead, built once in ``__init__``, would silently go stale across the vmap batch for the exact reason documented for ``ShallowWaterWeakForm``'s own fields. """ 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_taylor_basis from scimba_jax.linear_approximation.basis.general_bases import AnalyticBasis from scimba_jax.linear_approximation.galerkin.dg.flux import SIPGFlux from scimba_jax.linear_approximation.galerkin.dg.time_discrete_dg_scheme import ( TimeDiscreteDGscheme, ) from scimba_jax.linear_approximation.meshes.cartesian_mesh import cartesian_mesh from scimba_jax.linear_approximation.variables.variables_dg import VariablesDG 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_taylor_basis( y, i, m, order=k, out_dim=1 ), basis_type="scalar", ) return VariablesDG(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``, ``alpha`` a live pytree child (see the module docstring). ``bilinear_form`` writes ``alpha * grad(u).grad(v)`` directly, exactly as ``LaplacianWeakForm.bilinear_form`` writes the bare ``grad(u).grad(v)`` it specializes -- avoiding the ``A=alpha*I`` matmul ``EllipticWeakForm``'s generic path would otherwise trace. ``get_fields`` supplies ``fields["A"] = alpha * I`` on demand, read fresh from ``self.alpha`` at every call (not cached in ``__init__``), so that :class:`SIPGFlux`'s face terms stay consistent with the volume term for whichever ``alpha`` this particular vmap slot carries. """ 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 get_fields(self, pde_accessor=None): alpha, dim = self.alpha, self.dim dims = {"x": dim, "mu": 0} return { "A": self._wrap_static_field( lambda x: alpha * jnp.eye(dim), dims, lambda pytree, x, u_val=None: alpha * jnp.eye(dim), ) } 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 DG 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. """ h = 1.0 / n_cells sigma = (order + 1) * (order + 2) variables = make_variables(n_cells, order) weak_form = DiffusionWeakForm(dim=DIM, alpha=alpha) scheme = TimeDiscreteDGscheme( spatial_weak_form_factory=weak_form, variables=variables, flux=SIPGFlux(sigma=sigma, h=h), 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 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(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"DG, 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 DG 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="DG 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"DG, 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)