r"""1D parametric heat equation with an UNTRAINED DeepBasis, vmapped over runs. 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). Time-dependent counterpart of ``examples_jax/kernel/deepbasis_laplacian_ comparison.py``'s Part 1 -- a :class:`~scimba_jax.linear_approximation. basis.random_network_basis.DeepBasis` (a single MLP with ``n_basis`` outputs) is built with random, UNTRAINED weights and used as a fixed (ELM/random- feature-style) basis: only the linear coefficients ``dofsl`` are ever fit, via :class:`~scimba_jax.linear_approximation.collocation. time_discrete_collocation_scheme.TimeDiscreteCollocationScheme` (implicit Euler) driving, unmodified, a plain :class:`~scimba_jax.linear_approximation. collocation.collocation_elliptic.EllipticCollocationScheme` Newton solve at every time step -- exactly as it does for :class:`~scimba_jax.linear_approximation.basis.kernel_basis.KernelBasis` in ``examples_jax/kernel/time_dependent/1d_heat_equation.py``. Nothing in the scheme needed adapting: the RK-stage residual ``U_i + a_ii*dt*A(U_i) = ...`` is linear in the unknown coefficients ``dofsl`` regardless of how nonlinear the (frozen) basis functions are in ``x``, so Newton still converges in one step, and DeepBasis already exposes the ``all_values``/``kernel_function`` adapter :meth:`~scimba_jax.linear_approximation.variables. collocation_variables.CollocationVariables._classical_evaluate_vmap_pure` needs. Also borrows this module's parametric batching pattern from ``examples_jax/kernel/time_dependent/1d_heat_equation_parametric.py``: a physical parameter ``alpha`` is drawn ``N_RUNS_PARAMETRIC`` times and solved in one ``jax.vmap`` call, via :class:`~scimba_jax.physical_models. elliptic_pde.laplacians.ParametricDiffusionResidual` (the library counterpart of that file's local ``DiffusionResidual``). One thing is genuinely new here: each run ALSO draws its own PRNG key, so every vmap slot gets its own independent untrained DeepBasis, not just its own ``alpha`` -- otherwise "untrained random basis" would mean nothing (every run would share the exact same, arbitrarily-lucky-or-unlucky, random network). ``alpha`` and the basis key are therefore both vmapped inputs; unlike the Gaussian-RBF kernel basis (occasionally ill-conditioned for some point counts, per that module's docstring), a smooth (tanh) MLP basis was found to be well-behaved across every one of the 2000 random draws checked below -- no blow-ups, no NaNs. """ 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.random_network_basis import DeepBasis 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 ( ParametricDiffusionResidual, ) from scimba_jax.time_discrete.butcher_tableau import build_implicit_euler_tableau jax.config.update("jax_enable_x64", True) 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 # interior collocation points, also sets the DeepBasis width N_BASIS = 20 HIDDEN_SIZES = [20, 20] 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, n_basis, hidden_sizes, key): """Build an untrained DeepBasis collocation space on ``[0, 1]``. ``n_points`` sets both the number of interior collocation points and (to keep the two comparable, as ``examples_jax/kernel/deepbasis_laplacian_ comparison.py`` does) the DeepBasis width ``n_basis``. The network's weights come straight from ``key`` -- untrained, never optimized. """ x_colloc = jnp.linspace(0.0, 1.0, n_points).reshape(-1, 1) x_colloc_bc = jnp.array([[0.0], [1.0]]) domain = Segment1D(jnp.array([[0.0, 1.0]]), is_main_domain=True) basis = DeepBasis( dim=1, output_dim=1, n_basis=n_basis, hidden_sizes=hidden_sizes, key=key, basis_type="scalar", ) variables = CollocationVariables(basis=basis, nb_variables=1) return variables, domain, x_colloc, x_colloc_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) def _solve_one_draw( n_points: int, n_basis: int, hidden_sizes: list, dt: float, nt: int, x_plot: jnp.ndarray, alpha: jnp.ndarray, basis_key: jnp.ndarray, ) -> tuple[jnp.ndarray, jnp.ndarray, jnp.ndarray, jnp.ndarray]: """One collocation solve at diffusivity ``alpha`` with an untrained DeepBasis drawn from ``basis_key``, kept fully traceable so :func:`run_batch` can ``jax.vmap`` it over a whole batch of ``(alpha, basis_key)`` draws -- see ``1d_heat_equation_parametric.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_colloc, x_colloc_bc = make_variables( n_points, n_basis, hidden_sizes, basis_key ) def spatial_residual_factory(t): return ParametricDiffusionResidual( domain=domain, alpha=alpha, f_rhs=lambda x, mu: jnp.zeros(()) ) scheme = TimeDiscreteCollocationScheme( spatial_residual_factory=spatial_residual_factory, variables=variables, collocation_points=x_colloc, bc_collocation_points=x_colloc_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, n_basis: int, hidden_sizes: list, dt: float, nt: int, x_plot: jnp.ndarray, alpha_batch: jnp.ndarray, basis_key_batch: jnp.ndarray, ) -> tuple[jnp.ndarray, jnp.ndarray, jnp.ndarray, jnp.ndarray]: """Solves every ``(alpha, basis_key)`` draw 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, n_basis, hidden_sizes, dt, nt, x_plot ) return jax.vmap(solve_one)(alpha_batch, basis_key_batch) def compare_batching_time( n_points: int, n_basis: int, hidden_sizes: list, dt: float, nt: int, x_plot: jnp.ndarray, alpha_batch: jnp.ndarray, basis_key_batch: jnp.ndarray, ) -> tuple[jnp.ndarray, jnp.ndarray, jnp.ndarray, jnp.ndarray]: """Times 1 draw against the full ``N_RUNS_PARAMETRIC``-draw vmapped batch -- see ``1d_heat_equation_parametric.py``'s function of the same name for why an untimed warmup solve is needed first (absorbing the process's one-time XLA/pytree-registration warm-up so it doesn't artificially inflate whichever measurement runs first). 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, 4, hidden_sizes, T_FINAL / 2, 2, x_plot, alpha_batch[:1], basis_key_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, n_basis, hidden_sizes, dt, nt, x_plot, alpha_batch[:1], basis_key_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, n_basis, hidden_sizes, dt, nt, x_plot, alpha_batch, basis_key_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"each with its own untrained random DeepBasis, " f"n_basis={N_BASIS}, hidden_sizes={HIDDEN_SIZES}, 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, basis_key)`` 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 draws") 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="untrained DeepBasis, mean over draws") 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 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 draws") 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 draws") axes[1].legend(fontsize=8) axes[1].grid(True, alpha=0.3, which="both") fig.suptitle( "Untrained DeepBasis collocation, parametric heat equation, " f"alpha ~ Uniform([{ALPHA_MIN}, {ALPHA_MAX}]), " f"{N_RUNS_PARAMETRIC} independent (alpha, basis) draws, n_basis={N_BASIS}" ) fig.tight_layout() plt.show() if __name__ == "__main__": key = jax.random.PRNGKey(0) alpha_key, basis_key = jax.random.split(key) alpha_batch = jax.random.uniform( alpha_key, (N_RUNS_PARAMETRIC,), minval=ALPHA_MIN, maxval=ALPHA_MAX ) basis_key_batch = jax.random.split(basis_key, N_RUNS_PARAMETRIC) u_h_batch, u_ex_batch, abs_err_batch, rel_l2_err_batch = compare_batching_time( N_POINTS, N_BASIS, HIDDEN_SIZES, DT, NT, X_PLOT, alpha_batch, basis_key_batch ) print_summary(alpha_batch, rel_l2_err_batch) plot_results(X_PLOT, u_h_batch, u_ex_batch, abs_err_batch)