r"""Fixed (random) DeepBasis vs trained DeepBasis on a 2D heat equation. The PDE is: du/dt - alpha * (u_xx + u_yy) = 0 on [0, 1]^2, u = 0 on the boundary, u(x, y, 0) = sin(pi x) sin(pi y). Exact solution: u(x, y, t) = exp(-2 * alpha * pi^2 * t) sin(pi x) sin(pi y). """ import time import jax import jax.numpy as jnp import matplotlib.pyplot as plt from scimba_jax.domains.meshless_domains.domains_2d import Square2D 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.nonlinear_approximation.approximation_spaces.time_discrete_collocation_approximation_spaces import ( # noqa: E501 RKStageConsistencyResidual, TimeDiscreteCollocationApproximationSpace, ) from scimba_jax.nonlinear_approximation.integration.monte_carlo import ( DomainSampler, TensorizedSampler, ) from scimba_jax.nonlinear_approximation.numerical_solvers.projectors import Projector from scimba_jax.physical_models.abstract_physical_model import AbstractPhysicalModel 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) # ── Common parameters ────────────────────────────────────────────────────────── X_MIN, X_MAX = 0.0, 1.0 ALPHA = 1.2 FREQ = 2 T_FINAL = 0.01 NT = 50 DT = T_FINAL / NT N_BASIS = 20 HIDDEN_SIZES = [20] * 2 NEWTON_KWARGS = {"matrix_regularization": 1e-10} # Number of collocation points for the FIXED inner Newton solve (Parts 1 and 2): N_COLLOC_SOLVE = 200 N_BC_COLLOC_SOLVE = 150 # Number of collocation points used by the OUTER training loop (Part 2 only): N_COLLOC_TRAIN = 1500 N_EPOCHS = 250 # Residual weights; the two components are # [RK-stage consistency, initial-condition anchor]. RESIDUAL_WEIGHTS = {"interior": [1.0, 50.0]} def u0(xy: jnp.ndarray, mu=None) -> jnp.ndarray: """Initial condition u(x, y, 0) = sin(pi x) sin(pi y).""" return jnp.sin(FREQ * jnp.pi * xy[0:1]) * jnp.sin(FREQ * jnp.pi * xy[1:2]) def u_exact(xy: jnp.ndarray, t: float) -> jnp.ndarray: """Exact solution u(x, y, t) = exp(-2*alpha*pi^2*t) sin(pi x) sin(pi y).""" x, y = xy[:, 0:1], xy[:, 1:2] decay = jnp.exp(-2.0 * FREQ**2 * ALPHA * jnp.pi**2 * t) sin_x = jnp.sin(FREQ * jnp.pi * x) sin_y = jnp.sin(FREQ * jnp.pi * y) return (decay * sin_x * sin_y)[:, 0] domain = Square2D([(X_MIN, X_MAX), (X_MIN, X_MAX)], is_main_domain=True) collocation_key = jax.random.PRNGKey(1) _, _colloc_sample = TensorizedSampler([DomainSampler(domain)], bc=True).sample( collocation_key, N_COLLOC_SOLVE, N_BC_COLLOC_SOLVE ) X_COLLOC = _colloc_sample["interior"][0] X_COLLOC_BC = _colloc_sample["boundary"][0] def spatial_residual_factory(t: float) -> ParametricDiffusionResidual: """The (time-independent) spatial operator ``-alpha * (u_xx + u_yy)``.""" return ParametricDiffusionResidual( domain=domain, alpha=ALPHA, f_rhs=lambda x, mu: jnp.zeros(()) ) def build_scheme(variables: CollocationVariables) -> TimeDiscreteCollocationScheme: """Build a ``TimeDiscreteCollocationScheme`` sharing every setting but ``variables`` -- so Part 1 and Part 2 solve the exact same discretized problem, differing only in whether the DeepBasis is trained. """ return 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)), **NEWTON_KWARGS, ) key = jax.random.PRNGKey(0) # ═══════════════════════════════════════════════════════════════════════════════ # PART 1 -- fixed (random, untrained) DeepBasis: implicit-Euler time stepping # ═══════════════════════════════════════════════════════════════════════════════ print("=" * 60) print("PART 1 -- fixed DeepBasis (kernel solve, linear coefficients only)") print("=" * 60) key, subkey = jax.random.split(key) fixed_basis = DeepBasis( dim=2, output_dim=1, n_basis=N_BASIS, hidden_sizes=HIDDEN_SIZES, key=subkey, basis_type="scalar", ) variables_fixed = CollocationVariables(basis=fixed_basis, nb_variables=1) assembler_fixed = build_scheme(variables_fixed) t_start = time.perf_counter() dofsl_init_fixed = assembler_fixed.initialize(u0) dofsl_final_fixed, _ = assembler_fixed.solve(dofsl_init_fixed, t0=0.0, nt=NT) dofsl_final_fixed.block_until_ready() t_fixed = time.perf_counter() - t_start # ═══════════════════════════════════════════════════════════════════════════════ # PART 2 -- trained DeepBasis (Projector, self-supervised RK-stage + IC consistency) # ═══════════════════════════════════════════════════════════════════════════════ print() print("=" * 60) print("PART 2 -- trained DeepBasis (Projector, RK-stage + IC consistency)") print("=" * 60) key, subkey = jax.random.split(key) train_basis = DeepBasis( dim=2, output_dim=1, n_basis=N_BASIS, hidden_sizes=HIDDEN_SIZES, key=subkey, basis_type="scalar", ) variables_train = CollocationVariables(basis=train_basis, nb_variables=1) assembler_train = build_scheme(variables_train) space = TimeDiscreteCollocationApproximationSpace( dims={"x": 2, "dofsl": 1}, list_assemblers=[assembler_train], u0=u0, t0=0.0, nt=NT, model_type="x_dofsl", ) heat_pde_outer = AbstractPhysicalModel(main_domain=domain) heat_pde_outer.physical_residuals["interior"] = RKStageConsistencyResidual( domain, spatial_residual_factory(T_FINAL), a_ii=1.0, dt=DT, u0=u0 ) sampler_outer = TensorizedSampler([DomainSampler(domain)], bc=False) print("--- Training (SS-Broyden) ---") pinn = Projector( heat_pde_outer, space, sampler_outer, weights=RESIDUAL_WEIGHTS, optimizer="SS-Broyden", ) key, sample_dict = sampler_outer.sample(key, N_COLLOC_TRAIN) loss0 = pinn.evaluate_loss(space, sample_dict) print(f"Initial loss: {loss0:.6e}") print(f"Training ({N_EPOCHS} epochs) ...") t_start = time.perf_counter() key, pinn = pinn.project(key, space, N_EPOCHS, N_COLLOC_TRAIN) space = pinn.space loss_total = jnp.asarray(pinn.losses.losses_history["total"]).reshape(-1) jax.block_until_ready(loss_total) t_train = time.perf_counter() - t_start # ── Evaluation ──────────────────────────────────────────────────────────────── n_eval = 100 xs = jnp.linspace(X_MIN, X_MAX, n_eval) xx, yy = jnp.meshgrid(xs, xs) xy_eval = jnp.stack([xx.flatten(), yy.flatten()], axis=1) u_ref = u_exact(xy_eval, T_FINAL) u_fixed_eval = jax.vmap(lambda x: variables_fixed.evaluate(x))(xy_eval)[:, 0] u_i_fn, _u_n_fn, _u_init_fn = space.create_variables() dofsl_train_ = space.get_intermediate_values()[0] u_train = jax.vmap(u_i_fn, in_axes=(None, 0, None)) u_train_eval = u_train(space, xy_eval, dofsl_train_)[:, 0] err_fixed = jnp.abs(u_fixed_eval - u_ref) err_train = jnp.abs(u_train_eval - u_ref) l2_ref = float(jnp.sqrt(jnp.mean(u_ref**2))) l2_fixed = float(jnp.sqrt(jnp.mean(err_fixed**2))) / l2_ref l2_train = float(jnp.sqrt(jnp.mean(err_train**2))) / l2_ref print() print(f"{'':28s} {'L2 error':>12s} {'Time (s)':>10s} {'DOFs':>6s}") print(f"{'Fixed random DeepBasis':28s} {l2_fixed:12.4e} {t_fixed:10.3f} {N_BASIS:6d}") print(f"{'Trained DeepBasis':28s} {l2_train:12.4e} {t_train:10.3f} {N_BASIS:6d}") print( f"\nTrained/fixed time ratio: {t_train / t_fixed:.1f}x " f"(~{N_EPOCHS} epochs total, each ~1 full trajectory-solve-equivalent)" f"\nFixed/trained error ratio: {l2_fixed / l2_train:.1f}x" ) # ── Plots ───────────────────────────────────────────────────────────────────── fig, axs = plt.subplots(2, 3, figsize=(18, 10)) axs = axs.flatten() c0 = axs[0].contourf(xx, yy, u_ref.reshape(n_eval, n_eval), levels=50, cmap="turbo") fig.colorbar(c0, ax=axs[0]) axs[0].set_title(f"Exact solution at t={T_FINAL}") c1 = axs[1].contourf( xx, yy, u_fixed_eval.reshape(n_eval, n_eval), levels=50, cmap="turbo" ) fig.colorbar(c1, ax=axs[1]) axs[1].set_title(f"Fixed DeepBasis (L2={l2_fixed:.2e})") c2 = axs[2].contourf( xx, yy, u_train_eval.reshape(n_eval, n_eval), levels=50, cmap="turbo" ) fig.colorbar(c2, ax=axs[2]) axs[2].set_title(f"Trained DeepBasis (L2={l2_train:.2e})") err_fixed = err_fixed.reshape(n_eval, n_eval) / jnp.max(jnp.abs(u_ref)) c3 = axs[3].contourf(xx, yy, err_fixed, levels=50, cmap="turbo") fig.colorbar(c3, ax=axs[3]) axs[3].set_title("Error, fixed basis") err_train = err_train.reshape(n_eval, n_eval) / jnp.max(jnp.abs(u_ref)) c4 = axs[4].contourf(xx, yy, err_train, levels=50, cmap="turbo") fig.colorbar(c4, ax=axs[4]) axs[4].set_title("Error, trained basis") axs[5].semilogy(loss_total, "o-", markersize=3) axs[5].set_xlabel("Epoch") axs[5].set_ylabel("Loss") axs[5].set_title(f"Training loss history ({N_EPOCHS} SS-Broyden epochs)") axs[5].grid(True, alpha=0.3) for ax in axs[:5]: ax.set_xlabel("x") ax.set_ylabel("y") plt.tight_layout() plt.savefig("/tmp/deepbasis_heat2d_comparison.png", dpi=150) print("\nPlot saved to /tmp/deepbasis_heat2d_comparison.png") plt.show()