"""Benchmark : Classical RBF method vs Learnable RBF method. Problem : The PDE is: -Δu = f in Ω = [0,1]² with homogeneous Dirichlet BCs. Exact solution : u(x,y) = sin(π*x) * sin(π*y) Then f is computed from the PDE. """ 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.kernel_basis import KernelBasis from scimba_jax.linear_approximation.basis.kernel_function import ( GaussianKernel, GaussianKernelLocal, ) from scimba_jax.linear_approximation.collocation.collocation_elliptic import ( EllipticCollocationScheme, ) from scimba_jax.linear_approximation.variables.collocation_variables import ( CollocationVariables, ) from scimba_jax.nonlinear_approximation.approximation_spaces.collocation_approximation_spaces import ( CollocationEllipticApproximationSpace, ) 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.elliptic_pde.laplacians import ( LaplacianDirichletDG, LaplacianDirichletND, ) # ── Paramètres communs ──────────────────────────────────────────────────────── physical_dim = 2 out_dim = 1 seed = 1 newton_kwargs = {"max_iter": 3, "tol": 1e-12} # --- PDE definition and manufactured solution --- # Exact solution: u(x,y) = sin(π*x) * sin(π*y) def u_exact(xy: jnp.ndarray, mu=None) -> jnp.ndarray: x, y = xy[:, 0:1], xy[:, 1:2] return jnp.sin(jnp.pi * x) * jnp.sin(jnp.pi * y) # For -Δu = f, with u = sin(πx)sin(πy), we have Δu = -2π²u # So f = 2π²u = 2π² sin(πx) sin(πy) def f_rhs(xy: jnp.ndarray, mu=None) -> jnp.ndarray: """RHS for Laplacian: f = 2π² sin(πx) sin(πy)""" x, y = xy[0:1], xy[1:2] return 2.0 * jnp.pi**2 * jnp.sin(jnp.pi * x) * jnp.sin(jnp.pi * y) def dirichlet_bc(xy: jnp.ndarray, n, mu=None) -> jnp.ndarray: x, y = xy[0:1], xy[1:2] # noqa F841 return jnp.zeros_like(x) # Domain domain_x = [(0.0, 1.0), (0.0, 1.0)] dx = Square2D(domain_x, is_main_domain=True) # ── Classical Kernel method parameters ─────────────────────────────────────────────────── nb_centers = 100 # ── Learnable Kernel method parameters ────────────────────────────────────────────────── N_COLLOC = 300 N_BC_COLLOC = 100 N_EPOCHS = 80 # ═══════════════════════════════════════════════════════════════════════════════ # PARTIE 1 — Classical Kernel method (RBF with fixed centers and width) # ═══════════════════════════════════════════════════════════════════════════════ print("=" * 60) print("PART 1 — Classical Kernel method") print("=" * 60) # Create classical kernel basis with fixed centers and width # centers are uniformly spaced in the domain, and sigma is fixed (not dynamic) x_centers = jnp.linspace(0.01, 0.99, 9)[:, jnp.newaxis] y_centers = jnp.linspace(0.01, 0.99, 9)[:, jnp.newaxis] xy_centers = jnp.meshgrid(x_centers[:, 0], y_centers[:, 0]) xy_centers = jnp.stack(xy_centers, axis=-1).reshape(-1, 2) # Shape (nb_centers^2, 2) # plt.scatter(xy_centers[:, 0], xy_centers[:, 1], color="red", label="RBF centers (classical)") # plt.xlabel("x") # plt.ylabel("y") # plt.show() fixed_kernel = GaussianKernel(sigma=0.3, 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) pde_ref = LaplacianDirichletND( dx, lambda *args: f_rhs(*args), bc="weak", f_bc_rhs=lambda *args: dirichlet_bc(*args), ) # Sampling key = jax.random.PRNGKey(0) sampler = TensorizedSampler([DomainSampler(dx)], bc=True) key, sample_dict = sampler.sample(key, N_COLLOC, N_BC_COLLOC) x_centers = sample_dict["interior"][0] x_centers_bc = sample_dict["boundary"][0] assembler_classical = EllipticCollocationScheme( pde_ref, variables_classical, x_centers, x_centers_bc ) # solve the classical kernel method t0 = time.perf_counter() solved_classical = EllipticCollocationScheme.solve(assembler_classical, **newton_kwargs) solved_classical.variables.dofsl.block_until_ready() t_classical = time.perf_counter() - t0 # ═══════════════════════════════════════════════════════════════════════════════ # PARTIE 2 — RBF apprenable (Kernel optimisée par PINN) # ═══════════════════════════════════════════════════════════════════════════════ print() print("=" * 60) print("PARTIE 2 — RBF apprenable (Kernel + PINN)") print("=" * 60) learnable_kernel = GaussianKernelLocal(sigma=0.3, n_centers=81) learnable_basis = KernelBasis( dim=2, output_dim=1, kernel_function=learnable_kernel, centers=xy_centers, learnable_center_bool=False, learnable_kernel_bool=True, basis_type="scalar", ) learnable_variables = CollocationVariables(basis=learnable_basis) key = jax.random.PRNGKey(seed) key, subkey = jax.random.split(key) pde_learn = pde_ref # Same PDE assembler_learn = EllipticCollocationScheme( pde_learn, learnable_variables, x_centers, x_centers_bc ) space = CollocationEllipticApproximationSpace( dims={"x": physical_dim, "dofsl": 1}, list_assemblers=[assembler_learn], model_type="x_dofsl", newton_kwargs=newton_kwargs, ) model = LaplacianDirichletDG(main_domain=dx, f_rhs=f_rhs, bc="weak") sampler = TensorizedSampler([DomainSampler(dx)], bc=True) key, sample_dict = sampler.sample(key, N_COLLOC) pinn = Projector(model, space, sampler) loss0 = pinn.evaluate_loss(space, sample_dict) print(f"Loss initiale : {loss0:.6e}") print(f"Entraînement ({N_EPOCHS} époques) …") t0 = time.perf_counter() key, pinn = pinn.project(key, space, N_EPOCHS, N_COLLOC) new_loss = pinn.best_loss nspace = pinn.space loss_history = pinn.losses.losses_history jax.block_until_ready(jax.tree_util.tree_leaves(new_loss)) t_learnable = time.perf_counter() - t0 pinn.losses.loss_history = loss_history print(f"Loss finale : {new_loss['total']:.6e}") print(f"Temps entraînement (JIT inclus) : {t_learnable:.2f} s") # ── Évaluation et erreurs ───────────────────────────────────────────────────── n_eval = 50 x_plot = jnp.linspace(0.0, 1.0, n_eval)[:, jnp.newaxis] y_plot = jnp.linspace(0.0, 1.0, n_eval)[:, jnp.newaxis] xy_plot = jnp.meshgrid(x_plot[:, 0], y_plot[:, 0]) xy_plot = jnp.stack(xy_plot, axis=-1).reshape(-1, 2) # Shape (n_eval^2, 2) u_ref = u_exact(xy_plot)[:, 0] u_classical = jax.vmap(lambda x: solved_classical.variables.evaluate(x))(xy_plot)[:, 0] # Get initial DOFs BEFORE training (from the initial space) # get_intermediate_values() returns tuple with (stacked_dofsl,) # where stacked_dofsl has shape (n_assemblers, n_centers, n_vars) dofsl_list_init = space.get_intermediate_values() dofsl_list_final = nspace.get_intermediate_values() # Extract first (and only) assembler's DOFs: shape changes from (1, 10, 1) to (10, 1) dofsl_init = dofsl_list_init[0][0] dofsl_final = dofsl_list_final[0][0] (u_fn,) = space.create_variables() print("\nDOF shapes after extraction:") print(f" dofsl_init shape: {dofsl_init.shape}") print( f" dofsl_init min/max/mean: {jnp.min(dofsl_init):.4e} / {jnp.max(dofsl_init):.4e} / {jnp.mean(jnp.abs(dofsl_init)):.4e}" ) print(f" dofsl_final shape: {dofsl_final.shape}") print( f" dofsl_final min/max/mean: {jnp.min(dofsl_final):.4e} / {jnp.max(dofsl_final):.4e} / {jnp.mean(jnp.abs(dofsl_final)):.4e}" ) # For vmapped evaluation, we need to pass the stacked version to maintain compatibility # with the pre-processing functions that expect the assembler-stacked structure dofsl_init_stacked = dofsl_list_init[0] # Shape: (1, 10, 1) dofsl_final_stacked = dofsl_list_final[0] u_learn_init = jax.vmap(u_fn, in_axes=(None, 0, None))( space, xy_plot, dofsl_init_stacked )[:, 0] print(f"\nu_learn_init shape: {u_learn_init.shape}") print( f"u_learn_init min/max/mean: {jnp.min(u_learn_init):.4e} / {jnp.max(u_learn_init):.4e} / {jnp.mean(jnp.abs(u_learn_init)):.4e}" ) u_learn_final = jax.vmap(u_fn, in_axes=(None, 0, None))( nspace, xy_plot, dofsl_final_stacked )[:, 0] l2_classical = float(jnp.sqrt(jnp.mean((u_classical - u_ref) ** 2))) l2_learnable = float(jnp.sqrt(jnp.mean((u_learn_final - u_ref) ** 2))) print() print(f"{'':20s} {'Erreur L2':>12s} {'Temps (s)':>10s} {'DOFs':>6s}") # ── Résidu PINN ─────────────────────────────────────────────────────────────── def u_scalar(sp, x_1d, dofsl): return u_fn(sp, x_1d, dofsl)[0] def pinn_res(sp, x_1d, dofsl): H = jax.hessian(lambda x_: u_scalar(sp, x_, dofsl))(x_1d) lap = H[0, 0] + H[1, 1] return (-lap - f_rhs(x_1d)[0]) ** 2.0 res_init = jax.vmap(lambda x: pinn_res(space, x, dofsl_init_stacked))(xy_plot) res_final = jax.vmap(lambda x: pinn_res(nspace, x, dofsl_final_stacked))(xy_plot) print("\nResidue values:") print( f" res_init min/max/mean: {jnp.min(res_init):.4e} / {jnp.max(res_init):.4e} / {jnp.mean(res_init):.4e}" ) print( f" res_final min/max/mean: {jnp.min(res_final):.4e} / {jnp.max(res_final):.4e} / {jnp.mean(res_final):.4e}" ) # ── Plots ───────────────────────────────────────────────────────────────────── err_classical = jnp.abs(u_classical - u_ref) err_learn_final = jnp.abs(u_learn_final - u_ref) X, Y = jnp.meshgrid(x_plot[:, 0], y_plot[:, 0]) fig, axs = plt.subplots(2, 3, figsize=(18, 10)) axs = axs.flatten() c0 = axs[0].contourf(X, Y, u_ref.reshape(n_eval, n_eval), levels=50, cmap="turbo") fig.colorbar(c0, ax=axs[0]) axs[0].set_title("Exact Solution") axs[0].set_xlabel("x") axs[0].set_ylabel("y") c1 = axs[1].contourf(X, Y, u_classical.reshape(n_eval, n_eval), levels=50, cmap="turbo") fig.colorbar(c1, ax=axs[1]) axs[1].set_title(f"Classical (L2={l2_classical:.2e})") axs[1].set_xlabel("x") axs[1].set_ylabel("y") c2 = axs[2].contourf( X, Y, u_learn_final.reshape(n_eval, n_eval), levels=50, cmap="turbo" ) fig.colorbar(c2, ax=axs[2]) axs[2].set_title(f"Learnable (L2={l2_learnable:.2e})") axs[2].set_xlabel("x") axs[2].set_ylabel("y") c3 = axs[3].contourf( X, Y, err_classical.reshape(n_eval, n_eval), levels=50, cmap="turbo" ) fig.colorbar(c3, ax=axs[3]) axs[3].set_title("Error Classical") axs[3].set_xlabel("x") axs[3].set_ylabel("y") c4 = axs[4].contourf( X, Y, err_learn_final.reshape(n_eval, n_eval), levels=50, cmap="turbo" ) fig.colorbar(c4, ax=axs[4]) axs[4].set_title("Error Learnable") axs[4].set_xlabel("x") axs[4].set_ylabel("y") loss_total = jnp.asarray(loss_history["total"]).reshape(-1) axs[5].semilogy(loss_total, "o-", markersize=3) axs[5].set_xlabel("Epoch") axs[5].set_ylabel("Loss PINN") axs[5].set_title(f"Loss History ({N_EPOCHS} epochs)") axs[5].grid(True, alpha=0.3) plt.tight_layout() plt.show() # Show learned centers initial_centers = space.assemblers[0].variables.basis.centers learned_centers = nspace.assemblers[0].variables.basis.centers plt.figure(figsize=(6, 6)) plt.scatter( initial_centers[:, 0], initial_centers[:, 1], color="red", label="Initial Centers" ) plt.scatter( learned_centers[:, 0], learned_centers[:, 1], color="blue", label="Learned Centers", alpha=0.6, ) plt.xlabel("x") plt.ylabel("y") plt.title("RBF Centers: Initial vs Learned") plt.legend() plt.grid(True) plt.show() initial_sigma = space.assemblers[0].variables.basis.kernel_function.sigmas learned_sigma = nspace.assemblers[0].variables.basis.kernel_function.sigmas print(f"Sigma appris : {learned_sigma}") print(f"Sigma initial : {initial_sigma}")