# 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 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.integration.monte_carlo import ( DomainSampler, TensorizedSampler, ) from scimba_jax.physical_models.elliptic_pde.laplacians import 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.5, 0.5]) 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 = 6.0 r_border = base_radius + amplitude * jnp.cos(petals * theta) sdf = r - r_border return sdf[:, None] custom_domain = VolumetricDomain( domain_type="CustomUnitDisk", dim=2, sdf=CustomSDF(), bounds=[(0.0, 1.0), (0.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 = 6.0 center_bc = jnp.array([0.5, 0.5]) 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 custom_domain.add_bc_domain(bc) N_COLLOC = 1000 N_BC_COLLOC = 100 key = jax.random.PRNGKey(1) sampler = TensorizedSampler([DomainSampler(custom_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 # Exact solution: u(x,y) = sin(π*x) * sin(π*y) def u_exact(xy: jnp.ndarray, mu=None) -> jnp.ndarray: """Exact solution for batch of points (shape (n, 2))""" 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) -> 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 f_bc(xy: jnp.ndarray, n, mu) -> jnp.ndarray: """Dirichlet BC: u = 0 on boundary (homogeneous)""" return u_exact(xy, None) # PDE pde = LaplacianDirichletND( custom_domain, lambda *args: f_rhs(*args), bc="weak", f_bc_rhs=lambda *args: f_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, ) # Assemble and solve print("Assembling system...") A = scheme.assembly_scheme() print(f"Assembly matrix shape: {A.shape}") print("Solving...") scheme = EllipticCollocationScheme.solve(scheme) print("Done!") # Evaluate on test points and compute error n_eval = 100 key, sample_dict = sampler.sample(key, N_COLLOC, N_BC_COLLOC) xy_eval = sample_dict["interior"][0] u_exact_eval = jax.vmap(lambda x: u_exact(x, None))(xy_eval) u_computed = jax.vmap(lambda x: scheme.variables.evaluate(x))(xy_eval) print( f"\nEvaluation points shape: {xy_eval.shape}" f"\nExact solution shape: {u_exact_eval.shape}" f"\nComputed solution shape: {u_computed.shape}" ) u_exact_eval = jnp.squeeze(u_exact_eval) u_computed = jnp.squeeze(u_computed) error_l2 = jnp.linalg.norm(u_computed - u_exact_eval) / jnp.linalg.norm(u_exact_eval) print(f"\nRelative L2 error: {error_l2:.4e}") print(f"Max pointwise error: {jnp.max(jnp.abs(u_computed - u_exact_eval)):.4e}") # Plot # Build a regular grid, evaluate solutions on it and mask outside-domain points # grid resolution n_grid = 200 x_min, x_max = custom_domain.bounds[0] y_min, y_max = custom_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_exact_grid = jax.vmap(lambda x: u_exact(x, None))(grid_pts) u_comp_grid = jax.vmap(lambda x: scheme.variables.evaluate(x))(grid_pts) u_exact_grid = np.asarray(jnp.squeeze(u_exact_grid)).reshape(n_grid, n_grid) u_comp_grid = np.asarray(jnp.squeeze(u_comp_grid)).reshape(n_grid, n_grid) # mask outside points mask2d = mask_grid.reshape(n_grid, n_grid) u_exact_masked = np.ma.array(u_exact_grid, mask=mask2d) u_comp_masked = np.ma.array(u_comp_grid, mask=mask2d) error_grid = np.abs(u_comp_grid - u_exact_grid) error_masked = np.ma.array(error_grid, mask=mask2d) # Plot the three heatmaps side-by-side fig, axes = plt.subplots(1, 3, figsize=(15, 5), sharex=True, sharey=True) plots = [ (axes[0], u_exact_masked, "Exact Solution", "turbo", "u"), (axes[1], u_comp_masked, "Computed Solution", "turbo", "u"), (axes[2], error_masked, "Absolute Error", "turbo", "|u - u_exact|"), ] 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.savefig("/tmp/laplacian_solution.png", dpi=150) print("\nPlot saved to /tmp/laplacian_solution.png") plt.show()