r"""Classical (fixed) RBF kernel vs learnable kernel (trainable centers and variances) 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(FREQ*pi*x) sin(FREQ*pi*y). Exact solution: u(x, y, t) = exp(-2*FREQ^2*alpha*pi^2*t) sin(FREQ*pi*x) sin(FREQ*pi*y). Time-dependent counterpart of ``examples_jax/kernel/ 2d_poisson_learn_variance_centers.py``: same two kernel bases (a classical Gaussian ``KernelBasis`` with fixed centers/width vs. a ``KernelBasis`` built on a ``DeepFeatureGaussianKernel`` with trainable centers and per-feature variances), but driven through the time-stepping/training pattern of ``examples_jax/kernel/time_dependent/2d_heat_equation_deep_basis.py``: ``TimeDiscreteCollocationScheme`` + implicit Euler for the classical (untrained) part, ``TimeDiscreteCollocationApproximationSpace`` + ``RKStageConsistencyResidual`` + ``Projector`` for the trained part. """ import time import jax import jax.numpy as jnp import matplotlib.pyplot as plt from matplotlib.patches import Ellipse from scimba_jax.domains.meshless_domains.domains_2d import Square2D from scimba_jax.linear_approximation.basis.kernel_basis import KernelBasis from scimba_jax.linear_approximation.basis.kernel_function import ( DeepFeatureGaussianKernel, 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.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_CENTERS_1D = 5 # 5x5 = 25 RBF centers -- matches N_BASIS=20 in # 2d_heat_equation_deep_basis.py in cost: each training epoch backprops # through 50 unrolled implicit-Euler Newton solves, whose cost scales ~ # cubically with basis size, so more centers quickly dominates runtime # (measured: 81 centers -> ~1.15s/epoch, 25 centers -> ~0.3s/epoch). SIGMA = 0.3 NEWTON_KWARGS = {"matrix_regularization": 1e-10} # Number of collocation points for the FIXED inner Newton solve (both parts): 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(FREQ*pi*x) sin(FREQ*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*FREQ^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) # RBF centers: a fixed 9x9 grid, shared as the INITIAL centers of both the # classical (frozen) and the learnable (trained) kernel basis. x_c = jnp.linspace(0.05, 0.95, N_CENTERS_1D) xx_c, yy_c = jnp.meshgrid(x_c, x_c) xy_centers = jnp.stack([xx_c.flatten(), yy_c.flatten()], axis=-1) 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 kernel basis 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 -- classical kernel method (fixed centers, fixed width): implicit-Euler # time stepping, linear coefficients only # ═══════════════════════════════════════════════════════════════════════════════ print("=" * 60) print("PART 1 -- classical kernel method (fixed centers and width)") print("=" * 60) fixed_kernel = GaussianKernel(sigma=SIGMA, learnable_sigma=False) classical_basis = KernelBasis( dim=2, output_dim=1, kernel_function=fixed_kernel, centers=xy_centers, learnable_center_bool=False, basis_type="scalar", ) variables_classical = CollocationVariables(basis=classical_basis, nb_variables=1) assembler_classical = build_scheme(variables_classical) t_start = time.perf_counter() dofsl_init_classical = assembler_classical.initialize(u0) dofsl_final_classical, _ = assembler_classical.solve( dofsl_init_classical, t0=0.0, nt=NT ) dofsl_final_classical.block_until_ready() t_classical = time.perf_counter() - t_start # ═══════════════════════════════════════════════════════════════════════════════ # PART 2 -- learnable kernel (DeepFeatureGaussianKernel, trainable centers and # variances), trained via Projector, self-supervised RK-stage + IC consistency # ═══════════════════════════════════════════════════════════════════════════════ print() print("=" * 60) print("PART 2 -- learnable kernel (trainable centers + variances)") print("=" * 60) key, subkey = jax.random.split(key) learnable_kernel = DeepFeatureGaussianKernel(dim=2, sigma=SIGMA, key=subkey) learnable_basis = KernelBasis( dim=2, output_dim=1, kernel_function=learnable_kernel, centers=xy_centers, learnable_center_bool=True, learnable_kernel_bool=True, basis_type="scalar", ) variables_learn = CollocationVariables(basis=learnable_basis, nb_variables=1) assembler_learn = build_scheme(variables_learn) space = TimeDiscreteCollocationApproximationSpace( dims={"x": 2, "dofsl": 1}, list_assemblers=[assembler_learn], 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) nspace = 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_classical_eval = jax.vmap(lambda x: variables_classical.evaluate(x))(xy_eval)[:, 0] u_i_fn, _u_n_fn, _u_init_fn = nspace.create_variables() dofsl_train_ = nspace.get_intermediate_values()[0] u_train = jax.vmap(u_i_fn, in_axes=(None, 0, None)) u_learn_eval = u_train(nspace, xy_eval, dofsl_train_)[:, 0] err_classical = jnp.abs(u_classical_eval - u_ref) err_learn = jnp.abs(u_learn_eval - u_ref) l2_ref = float(jnp.sqrt(jnp.mean(u_ref**2))) l2_classical = float(jnp.sqrt(jnp.mean(err_classical**2))) / l2_ref l2_learn = float(jnp.sqrt(jnp.mean(err_learn**2))) / l2_ref print() print(f"{'':28s} {'L2 error':>12s} {'Time (s)':>10s} {'Centers':>8s}") print( f"{'Classical kernel':28s} {l2_classical:12.4e} " f"{t_classical:10.3f} {xy_centers.shape[0]:8d}" ) print( f"{'Learnable kernel':28s} {l2_learn:12.4e} " f"{t_train:10.3f} {xy_centers.shape[0]:8d}" ) print( f"\nTrained/classical time ratio: {t_train / t_classical:.1f}x " f"(~{N_EPOCHS} epochs total, each ~1 full trajectory-solve-equivalent)" f"\nClassical/trained error ratio: {l2_classical / l2_learn:.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_classical_eval.reshape(n_eval, n_eval), levels=50, cmap="turbo" ) fig.colorbar(c1, ax=axs[1]) axs[1].set_title(f"Classical kernel (L2={l2_classical:.2e})") c2 = axs[2].contourf( xx, yy, u_learn_eval.reshape(n_eval, n_eval), levels=50, cmap="turbo" ) fig.colorbar(c2, ax=axs[2]) axs[2].set_title(f"Learnable kernel (L2={l2_learn:.2e})") err_classical_plot = err_classical.reshape(n_eval, n_eval) / jnp.max(jnp.abs(u_ref)) c3 = axs[3].contourf(xx, yy, err_classical_plot, levels=50, cmap="turbo") fig.colorbar(c3, ax=axs[3]) axs[3].set_title("Error, classical kernel") err_learn_plot = err_learn.reshape(n_eval, n_eval) / jnp.max(jnp.abs(u_ref)) c4 = axs[4].contourf(xx, yy, err_learn_plot, levels=50, cmap="turbo") fig.colorbar(c4, ax=axs[4]) axs[4].set_title("Error, learnable kernel") 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/kernel_heat2d_learn_variance_centers.png", dpi=150) print("\nPlot saved to /tmp/kernel_heat2d_learn_variance_centers.png") plt.show() # ── Learned centers and kernel footprints ────────────────────────────────────── def kernel_footprint_ellipses(kernel_function, centers, max_ellipses=25): """Local Mahalanobis footprint (level ``exp(-1)``) of a DeepFeatureGaussianKernel. Linearizing the feature map ``phi`` around each center ``y_i`` gives ``K(x, y_i) ~= exp(-(x - y_i)^T A_i (x - y_i))`` with ``A_i = J_i^T diag(exp(log_weights)) J_i``, ``J_i = jacobian(phi)(y_i)``. The ellipse boundary ``(x - y_i)^T A_i (x - y_i) = 1`` is what a classical isotropic RBF would draw as its "1-sigma" circle. Only a bounded, evenly-spaced subset of centers gets an ellipse: drawing one per center (there can be dozens) quickly turns the figure into an unreadable smear of overlapping outlines. Args: kernel_function: A ``DeepFeatureGaussianKernel`` instance. centers: Center points, shape ``(nb_centers, dim)``. max_ellipses: Maximum number of ellipses to draw. Returns: Tuple ``(sub_centers, semi_axes, angles_deg)``. """ weights = jnp.exp(kernel_function.log_weights) stride = max(1, centers.shape[0] // max_ellipses) sub_centers = centers[::stride] def ellipse_params(y): jac = jax.jacobian(lambda x: x + kernel_function.feature_map(x))(y) quad_form = jac.T @ jnp.diag(weights) @ jac eigvals, eigvecs = jnp.linalg.eigh(quad_form) semi_axes = 1.0 / jnp.sqrt(jnp.clip(eigvals, 1e-8, None)) angle = jnp.degrees(jnp.arctan2(eigvecs[1, -1], eigvecs[0, -1])) return semi_axes, angle semi_axes, angles = jax.vmap(ellipse_params)(sub_centers) return sub_centers, semi_axes, angles fig, ax = plt.subplots(figsize=(6, 6)) initial_centers = xy_centers learned_centers = nspace.assemblers[0].variables.basis.centers learned_kernel = nspace.assemblers[0].variables.basis.kernel_function for centers, kernel_fn, color, label in [ (initial_centers, learnable_kernel, "red", "Initial"), (learned_centers, learned_kernel, "blue", "Learned"), ]: sub_centers, semi_axes, angles = kernel_footprint_ellipses(kernel_fn, centers) for c, (a, b), angle in zip(sub_centers, semi_axes, angles): ax.add_patch( Ellipse( (c[0], c[1]), width=2 * a, height=2 * b, angle=angle, edgecolor=color, facecolor="none", alpha=0.25, linewidth=0.8, ) ) ax.scatter( centers[:, 0], centers[:, 1], color=color, s=15, label=f"{label} Centers" ) ax.set_xlabel("x") ax.set_ylabel("y") ax.set_title( "Kernel Centers: Initial vs Learned\n(faint ellipses: kernel footprint, subsampled)" ) ax.legend() ax.grid(True) ax.set_aspect("equal") plt.show() print(f"log_weights initial : {learnable_kernel.log_weights}") print(f"log_weights appris : {learned_kernel.log_weights}")