# import necessary libraries import jax import jax.numpy as jnp import matplotlib.pyplot as plt import numpy as np from scimba_jax.domains.domain_mapping import DomainMapping from scimba_jax.domains.meshless_domains.base import SurfacicDomain, VolumetricDomain from scimba_jax.domains.meshless_domains.domains_1d import Segment1D from scimba_jax.domains.sdf import SignedDistance 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.mapping.mapping import InvertibleFunction 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, ) # Create custome flower domain class Id(InvertibleFunction): def __init__(self): super().__init__(f=lambda x: x, f_inv=lambda y: y) class CustomSDF(SignedDistance): """A custom SDF class that defines the signed distance function for unit 2D disk.""" def __init__(self): super().__init__( dim=2, threshold=0.0, # threshold for inside points (sdf < -threshold) ) def __call__(self, pts: jnp.ndarray) -> jnp.ndarray: """Compute the signed distance function for a flower-shaped domain. Use polar coordinates: r - r_border(theta) where r_border(theta) = 1 + amplitude * cos(petals * theta). This produces a smooth petal-shaped boundary and avoids singular divisions present in the previous formula. """ # pts: (N, dim) tensor x, y = pts[:, 0], pts[:, 1] # center and scale to fit inside unit square center = jnp.array([0.0, 0.0]) x_rel = x - center[0] y_rel = y - center[1] r = jnp.sqrt(x_rel**2 + y_rel**2) theta = jnp.arctan2(y_rel, x_rel) base_radius = 0.35 amplitude = 0.15 petals = 3.0 r_border = base_radius + amplitude * jnp.cos(petals * theta) sdf = r - r_border return sdf[:, None] flower_domain = VolumetricDomain( domain_type="CustomUnitDisk", dim=2, sdf=CustomSDF(), bounds=[(-1.0, 1.0), (-1.0, 1.0)], is_main_domain=True, ) # boundary mapping parameters must match SDF center/scale base_radius_bc = 0.35 amplitude_bc = 0.15 petals_bc = 3.0 center_bc = jnp.array([0.0, 0.0]) def flower_map(theta: jnp.ndarray) -> jnp.ndarray: th = theta[..., 0] r = base_radius_bc + amplitude_bc * jnp.cos(petals_bc * th) x = center_bc[0] + r * jnp.cos(th) y = center_bc[1] + r * jnp.sin(th) return jnp.stack([x, y], axis=-1) def flower_jac(theta: jnp.ndarray) -> jnp.ndarray: th = theta[..., 0] r = base_radius_bc + amplitude_bc * jnp.cos(petals_bc * th) dr_dth = -amplitude_bc * petals_bc * jnp.sin(petals_bc * th) dx_dth = dr_dth * jnp.cos(th) - r * jnp.sin(th) dy_dth = dr_dth * jnp.sin(th) + r * jnp.cos(th) # return shape (..., 2, 1) return jnp.stack([dx_dth, dy_dth], axis=-1)[..., None] surface_mapping = DomainMapping(1, 2, flower_map, flower_jac, None, None) bc = SurfacicDomain( domain_type="CustomFlowerBoundary", parametric_domain=Segment1D((0, 2 * jnp.pi)), surface=surface_mapping, ) # we then add the boundary domain to our custom domain flower_domain.add_bc_domain(bc) N_EPOCHS = 100 N_COLLOC = 300 N_BC_COLLOC = 100 newton_kwargs = {"max_iter": 3, "tol": 1e-12} key = jax.random.PRNGKey(1) sampler = TensorizedSampler([DomainSampler(flower_domain)], bc=True) key, sample_dict = sampler.sample(key, N_COLLOC, N_BC_COLLOC) xy_centers = sample_dict["interior"][0] xy_centers_bc = sample_dict["boundary"][0] # show domain plt.figure(figsize=(6, 6)) plt.scatter( xy_centers[:, 0], xy_centers[:, 1], color="blue", label="Interior Points", s=8 ) plt.scatter( xy_centers_bc[:, 0], xy_centers_bc[:, 1], color="red", label="Boundary Points", s=12 ) plt.title("Sampled Points in flower Domain") plt.xlabel("x") plt.ylabel("y") plt.legend() plt.axis("equal") plt.grid() plt.show() # Define PDE, basis, and scheme """Problem : The PDE is: -Δu = 1 in Ω = [0,1]² with homogeneous Dirichlet BCs. No exact solution known. """ def f_rhs(xy: jnp.ndarray, mu=None) -> jnp.ndarray: """RHS for Helmholtz: f = 1""" x, y = xy[0:1], xy[1:2] # noqa F841 return jnp.ones_like(x) 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) # PDE pde = LaplacianDirichletND( flower_domain, lambda *args: f_rhs(*args), bc="weak", f_bc_rhs=lambda *args: dirichlet_bc(*args), ) # Basis basis = KernelBasis( dim=2, output_dim=1, kernel_function=GaussianKernel(sigma=0.5), centers=xy_centers, learnable_center_bool=False, basis_type="scalar", ) variables = CollocationVariables(basis=basis, nb_variables=1) # Scheme scheme = EllipticCollocationScheme( pde=pde, variables=variables, collocation_points=xy_centers, bc_collocation_points=xy_centers_bc, ) solved_scheme = EllipticCollocationScheme.solve(scheme, **newton_kwargs) print("\nScheme solved.") #################################################################### # Learnable kernel part # Use a small, independent set of centers for the dynamic basis: reusing # all N_COLLOC=300 collocation points as centers would make nb_centers=300 # (instead of the intended 81), which multiplies the cost of the per-epoch # dense Jacobian/Newton solve ~4x and mismatches GaussianKernelLocal's # n_centers=81 sigma array. N_CENTERS = 9**2 key, center_sample_dict = sampler.sample(key, N_CENTERS) learnable_centers = center_sample_dict["interior"][0] learnable_kernel = GaussianKernelLocal(sigma=0.5, n_centers=N_CENTERS) learnable_basis = KernelBasis( dim=2, output_dim=1, kernel_function=learnable_kernel, centers=learnable_centers, learnable_center_bool=True, learnable_kernel_bool=True, basis_type="scalar", ) learnable_variables = CollocationVariables(basis=learnable_basis) assembler_learn = EllipticCollocationScheme( pde=pde, variables=learnable_variables, collocation_points=xy_centers, bc_collocation_points=xy_centers_bc, ) space = CollocationEllipticApproximationSpace( dims={"x": 2, "dofsl": 1}, list_assemblers=[assembler_learn], model_type="x_dofsl", newton_kwargs=newton_kwargs, ) model = LaplacianDirichletDG(main_domain=flower_domain, f_rhs=f_rhs, bc="weak") sampler = TensorizedSampler([DomainSampler(flower_domain)], bc=True) key, sample_dict = sampler.sample(key, N_COLLOC) pinn = Projector(model, space, sampler) # Train the model 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)) # get the model dofsl_list_final = nspace.get_intermediate_values() dofsl_final_stacked = dofsl_list_final[0] (u_fn,) = space.create_variables() # Evaluate on test points and compute error n_eval = 1_000 key, sample_dict = sampler.sample(key, n_eval, N_BC_COLLOC) xy_eval = sample_dict["interior"][0] u_computed = jax.vmap(lambda x: solved_scheme.variables.evaluate(x))(xy_eval)[:, 0] # ── Residuals ─────────────────────────────────────────────────────────────── # Classical residual def u_scalar(x_test): # Single point (2,) -> scalar; required by jax.hessian return solved_scheme.variables.evaluate(x_test)[0] def classical_res(x_test): H = jax.hessian(u_scalar)(x_test) lap = H[0, 0] + H[1, 1] return (-lap - f_rhs(x_test)[0]) ** 2.0 # PINN residual def u_scalar_pinn(sp, x_test, dofsl): return u_fn(sp, x_test, dofsl)[0] def pinn_res(sp, x_test, dofsl): H = jax.hessian(lambda x_: u_scalar_pinn(sp, x_, dofsl))(x_test) lap = H[0, 0] + H[1, 1] return (-lap - f_rhs(x_test)[0]) ** 2.0 # ── Plots ───────────────────────────────────────────────────────────────────── # Plot 1: The computed solution u on the domain, masking outside points # Build a regular grid, evaluate solutions on it and mask outside-domain points # grid resolution n_grid = 200 x_min, x_max = flower_domain.bounds[0] y_min, y_max = flower_domain.bounds[1] xi = np.linspace(x_min, x_max, n_grid) yi = np.linspace(y_min, y_max, n_grid) xx, yy = np.meshgrid(xi, yi) grid_pts = jnp.array(np.stack([xx.ravel(), yy.ravel()], axis=1)) # evaluate SDF on grid to create mask (True outside) sdf_vals = np.asarray(CustomSDF()(grid_pts))[:, 0] mask_grid = sdf_vals > 0 # evaluate exact and computed solutions on the grid u_comp_grid = jax.vmap(lambda x: solved_scheme.variables.evaluate(x))(grid_pts) u_comp_grid = np.asarray(jnp.squeeze(u_comp_grid)).reshape(n_grid, n_grid) u_pinn_grid = jax.vmap(lambda x: u_fn(nspace, x, dofsl_final_stacked))(grid_pts) u_pinn_grid = np.asarray(jnp.squeeze(u_pinn_grid)).reshape(n_grid, n_grid) # mask outside points mask2d = mask_grid.reshape(n_grid, n_grid) u_comp_masked = np.ma.array(u_comp_grid, mask=mask2d) u_pinn_masked = np.ma.array(u_pinn_grid, mask=mask2d) # residual residual_grid = jax.vmap(classical_res)(grid_pts) residual_grid = np.asarray(residual_grid).reshape(n_grid, n_grid) residual_masked = np.ma.array(residual_grid, mask=mask2d) residual_pinn_grid = jax.vmap(lambda x: pinn_res(nspace, x, dofsl_final_stacked))( grid_pts ) residual_pinn_grid = np.asarray(residual_pinn_grid).reshape(n_grid, n_grid) residual_pinn_masked = np.ma.array(residual_pinn_grid, mask=mask2d) # Plot the three heatmaps side-by-side with the domain boundary overlaid # and the residual below the solution fig, axes = plt.subplots(2, 2, figsize=(10, 6), sharex=True, sharey=True) plots = [ (axes[0, 0], u_comp_masked, "Computed Solution", "turbo", "u"), (axes[0, 1], u_pinn_masked, "PINN Solution", "turbo", "u"), (axes[1, 0], residual_masked, "Residual", "turbo", "Residual"), (axes[1, 1], residual_pinn_masked, "PINN Residual", "turbo", "PINN Residual"), ] for ax, values, title, cmap, cbar_label in plots: heatmap = ax.imshow( values, extent=(x_min, x_max, y_min, y_max), origin="lower", cmap=cmap ) fig.colorbar(heatmap, ax=ax, label=cbar_label) # contour_line = ax.contour(xx, yy, sdf_vals.reshape(n_grid, n_grid), levels=[0], colors=['purple'], linewidths=2, linestyles='dashed') ax.set_title(title) ax.set_xlabel("x") ax.set_ylabel("y") ax.set_aspect("equal", adjustable="box") plt.tight_layout() plt.show()