# 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.nonlinear_approximation.integration.monte_carlo import ( DomainSampler, TensorizedSampler, ) from scimba_jax.physical_models.elliptic_pde.laplacians import LaplacianDirichletND # Batman level-set helpers, written in JAX-friendly form. def batman_upper(x): x_abs = jnp.abs(x) conditions = [ x_abs < 0.5, (x_abs >= 0.5) & (x_abs < 0.75), (x_abs >= 0.75) & (x_abs < 1.0), (x_abs >= 1.0) & (x_abs <= 7.0), ] functions = [ lambda x: jnp.full_like(x, 2.25), lambda x: 3.0 * x + 0.75, lambda x: 9.0 - 8.0 * x, lambda x: 3.0 * jnp.sqrt(jnp.maximum(0.0, 1.0 - (x / 7.0) ** 2)), ] return jnp.select(conditions, [f(x_abs) for f in functions], default=0.0) def batman_lower(x): x_abs = jnp.abs(x) conditions = [ x_abs >= 4.0, x_abs < 4.0, ] functions = [ lambda x: -3.0 * jnp.sqrt(jnp.maximum(0.0, 1.0 - (x / 7.0) ** 2)), lambda x: ( (x / 2.0) - 0.09137 * (x**2) - 3.0 + jnp.sqrt(jnp.maximum(0.0, 1.0 - (jnp.abs(x - 2.0) - 1.0) ** 2)) ), ] return jnp.select(conditions, [f(x_abs) for f in functions], default=0.0) class BatmanSDF(SignedDistance): """Signed-distance-like level set for the Batman shape.""" def __init__(self): super().__init__(dim=2, threshold=0.0) def __call__(self, pts: jnp.ndarray) -> jnp.ndarray: pts = jnp.atleast_2d(pts) x = pts[:, 0] y = pts[:, 1] f_upper = y - batman_upper(x) f_lower = batman_lower(x) - y f = jnp.maximum(f_upper, f_lower) return f[:, None] custom_domain = VolumetricDomain( domain_type="CustomBatman", dim=2, sdf=BatmanSDF(), bounds=[(-7.2, 7.2), (-3.5, 3.5)], is_main_domain=True, ) def batman_boundary_map(theta: jnp.ndarray) -> jnp.ndarray: t = theta[0] def top_branch(tt): x = -7.0 + 28.0 * tt y = batman_upper(x) return jnp.stack([x, y]) def bottom_branch(tt): x = 7.0 - 28.0 * (tt - 0.5) y = batman_lower(x) return jnp.stack([x, y]) return jax.lax.cond(t < 0.5, top_branch, bottom_branch, t) surface_mapping = DomainMapping(1, 2, batman_boundary_map, None, None, None) bc = SurfacicDomain( domain_type="CustomBatmanBoundary", parametric_domain=Segment1D((0.0, 1.0)), surface=surface_mapping, ) custom_domain.add_bc_domain(bc) N_COLLOC = 1000 N_BC_COLLOC = 300 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] 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 Batman Domain") plt.xlabel("x") plt.ylabel("y") plt.legend() plt.axis("equal") plt.grid() if "agg" not in plt.get_backend().lower(): plt.show() else: plt.close() # Define PDE, basis, and scheme def u_exact(xy: jnp.ndarray, mu=None) -> jnp.ndarray: """Exact solution for batch of points (shape (n, 2))""" # xy = jnp.atleast_2d(xy) x, y = xy[0:1], xy[1:2] return jnp.sin(jnp.pi * x) * jnp.sin(jnp.pi * y) def f_rhs(xy: jnp.ndarray, mu) -> jnp.ndarray: """RHS for Laplacian: f = 2π² sin(πx) sin(πy)""" # xy = jnp.atleast_2d(xy) 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 set to the exact solution on the Batman boundary.""" # xy = jnp.atleast_2d(xy) return u_exact(xy, None) pde = LaplacianDirichletND( custom_domain, lambda *args: f_rhs(*args), bc="weak", f_bc_rhs=lambda *args: f_bc(*args), ) 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 = EllipticCollocationScheme( pde=pde, variables=variables, collocation_points=xy_centers, bc_collocation_points=xy_centers_bc, ) 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 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 on a regular grid, masked by the Batman level set 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)) sdf_vals = np.asarray(BatmanSDF()(grid_pts))[:, 0] mask_grid = sdf_vals > 0 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) 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) x_boundary = np.linspace(-7.0, 7.0, 500) y_top = np.asarray(batman_upper(jnp.array(x_boundary))) y_bottom = np.asarray(batman_lower(jnp.array(x_boundary))) cmap_exact = plt.cm.turbo.copy() cmap_exact.set_bad(alpha=0.0) cmap_error = plt.cm.magma.copy() cmap_error.set_bad(alpha=0.0) fig, axes = plt.subplots(1, 3, figsize=(15, 5), sharex=True, sharey=True) plots = [ (axes[0], u_exact_masked, "Exact Solution", cmap_exact, "u"), (axes[1], u_comp_masked, "Computed Solution", cmap_exact, "u"), (axes[2], error_masked, "Absolute Error", cmap_error, "|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, interpolation="nearest", ) fig.colorbar(heatmap, ax=ax, label=cbar_label) ax.plot(x_boundary, y_top, color="white", linewidth=1.2) ax.plot(x_boundary, y_bottom, color="white", linewidth=1.2) ax.set_title(title) ax.set_xlabel("x") ax.set_ylabel("y") ax.set_aspect("equal", adjustable="box") plt.tight_layout() plt.show()