"""Solve a 2D Helmholtz equation using kernel-based collocation. The PDE is given by: -∇²u + k²u = f in Ω = [-1,1]² u = g on ∂Ω Manufactured solution: u(x,y) = sin(π*x)*sin(π*y) with: k = 2 (wave number) g = 0 (homogeneous Dirichlet BC) RHS f is computed from the PDE: f = -∇²u + k²u """ import jax import jax.numpy as jnp import matplotlib.pyplot as plt from scimba_jax.domains.meshless_domains.base import VolumetricDomain 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 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.nonlinear_approximation.model_class.funcparam_vectorial import ( ParamScalarFunction, ) from scimba_jax.physical_models.abstract_physical_model import ( PHYSICAL_RESIDUALS_TYPE, AbstractPhysicalModel, ) from scimba_jax.physical_models.abstract_residuals import ( NDARRAYS_FUNC_TYPE, PARAM_FUNC_TYPE, InteriorResidual, ) from scimba_jax.physical_models.boundary_residuals import DirichletResidual N_COLLOC = 1_000 N_BC_COLLOC = 200 key = jax.random.PRNGKey(0) # Helmholtz wave number k = 2.0 # ============================================================================ # Manufactured Solution and Problem Definition # ============================================================================ def u_exact(x: jnp.ndarray, mu=None) -> jnp.ndarray: """Exact solution: u(x,y) = sin(π*x)*sin(π*y)""" x_coord, y_coord = x[0:1], x[1:2] return (jnp.sin(jnp.pi * x_coord) * jnp.sin(jnp.pi * y_coord)).squeeze() def d_fn(x: jnp.ndarray, mu=None) -> jnp.ndarray: """Diffusion coefficient: d(x,y) = I""" return jnp.eye(2) # Identity diffusion in 2D def c_fn(x: jnp.ndarray, mu=None) -> jnp.ndarray: """Reaction coefficient: c = k²""" return k**2 def f_rhs(x: jnp.ndarray, mu=None) -> jnp.ndarray: """ RHS for: -∇²u + k²u = f With u = sin(π*x)*sin(π*y): ∂²u/∂x² = -π²*sin(π*x)*sin(π*y) ∂²u/∂y² = -π²*sin(π*x)*sin(π*y) ∇²u = -2π²*sin(π*x)*sin(π*y) f = -∇²u + k²u = 2π²*sin(π*x)*sin(π*y) + k²*sin(π*x)*sin(π*y) = (2π² + k²)*sin(π*x)*sin(π*y) """ x_coord, y_coord = x[0:1], x[1:2] u = jnp.sin(jnp.pi * x_coord) * jnp.sin(jnp.pi * y_coord) laplacian_u = -2.0 * jnp.pi**2 * u f = -laplacian_u + k**2 * u return f def f_bc(x: jnp.ndarray, n, mu=None) -> jnp.ndarray: """Boundary condition: u = 0 (homogeneous Dirichlet)""" x_coord, _ = x[0:1], x[1:2] return jnp.zeros_like(x_coord) # ============================================================================ # Verify Manufactured Solution # ============================================================================ print(f"Verifying manufactured solution (2D Helmholtz with k={k})...") # Test at a single point xy_test = jnp.array([0.3, 0.5]) # Compute exact solution and its derivatives u_val = u_exact(xy_test) print(f"u(0.3, 0.5) = {u_val:.6f}") # Compute derivatives manually x, y = 0.3, 0.5 u_exact_val = jnp.sin(jnp.pi * x) * jnp.sin(jnp.pi * y) d2u_dx2_exact = -(jnp.pi**2) * jnp.sin(jnp.pi * x) * jnp.sin(jnp.pi * y) d2u_dy2_exact = -(jnp.pi**2) * jnp.sin(jnp.pi * x) * jnp.sin(jnp.pi * y) laplacian_exact = d2u_dx2_exact + d2u_dy2_exact # Verify PDE: -∇²u + k²u = f residual_val = -laplacian_exact + k**2 * u_exact_val f_val_raw = f_rhs(xy_test) f_val = float(jnp.squeeze(f_val_raw)) print(f"Residual error (should be 0): {abs(f_val - residual_val):.6e}") # ============================================================================ # Setup Domain and Sampling # ============================================================================ print("\nSetting up 2D domain and sampling...") domain_xy = [(-1.0, 1.0), (-1.0, 1.0)] dx = Square2D(domain_xy, is_main_domain=True) sampler = TensorizedSampler([DomainSampler(dx)], 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] print(f"Interior collocation points: {xy_centers.shape}") print(f"Boundary collocation points: {xy_centers_bc.shape}") # ============================================================================ # Setup Basis and Variables # ============================================================================ 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) # ============================================================================ # Create Helmholtz PDE model # ============================================================================ class HelmholtzResidual(InteriorResidual): """Custom residual for -∇²u + k²u = f in 2D.""" def __init__( self, domain: VolumetricDomain, f_rhs: NDARRAYS_FUNC_TYPE | None = None, d_fn: NDARRAYS_FUNC_TYPE | None = None, c_fn: NDARRAYS_FUNC_TYPE | None = None, ): super().__init__(domain=domain, size=1, model_type="x_mu", f_rhs=f_rhs) self.d = d_fn self.c = c_fn def construct_residual(self, *vars: PARAM_FUNC_TYPE) -> PARAM_FUNC_TYPE: rho = vars[0] assert isinstance(rho, ParamScalarFunction) # Laplacian: -∇²u residual = rho.anisotropic_laplacian("x", self.d) # Add reaction term: k²u if self.c is not None: c_times_u = rho * self.c residual = residual + c_times_u return residual class HelmholtzND(AbstractPhysicalModel): """A 2D Helmholtz equation.""" def __init__( self, main_domain: VolumetricDomain, f_rhs: NDARRAYS_FUNC_TYPE | None = None, bc: str = "weak", f_bc_rhs: NDARRAYS_FUNC_TYPE | None = None, d: NDARRAYS_FUNC_TYPE | None = None, c: NDARRAYS_FUNC_TYPE | None = None, ): super().__init__(main_domain=main_domain) self.physical_residuals: PHYSICAL_RESIDUALS_TYPE = { self.main_domain.get_label(): HelmholtzResidual( domain=main_domain, f_rhs=f_rhs, d_fn=d, c_fn=c, ), } if bc == "weak": for boundary in self.boundaries: self.physical_residuals[boundary] = DirichletResidual( domain=self.boundaries[boundary], model_type="x_mu", f_rhs=f_bc_rhs, ) pde = HelmholtzND( main_domain=dx, f_rhs=f_rhs, bc="weak", f_bc_rhs=f_bc, d=d_fn, c=c_fn, ) scheme = EllipticCollocationScheme( pde=pde, variables=variables, collocation_points=xy_centers, bc_collocation_points=xy_centers_bc, ) # ============================================================================ # Solve # ============================================================================ scheme.variables.use_scan = True print("\nAssembling system (Helmholtz 2D)...") A_mat = scheme.assembly_scheme() print(f"Assembly matrix shape: {A_mat.shape}") print("Solving system (Helmholtz 2D)...") scheme = EllipticCollocationScheme.solve(scheme, max_iter=2) print("Done!") # ============================================================================ # Evaluate and Compare # ============================================================================ n_eval = 50 # 50x50 evaluation grid x_eval = jnp.linspace(-1, 1, n_eval) y_eval = jnp.linspace(-1, 1, n_eval) xx, yy = jnp.meshgrid(x_eval, y_eval) xy_eval = jnp.stack([xx.flatten(), yy.flatten()], axis=-1) print(f"\nEvaluating solutions on {n_eval}x{n_eval} grid...") u_exact_eval = jax.vmap(u_exact)(xy_eval) # Vmap over evaluation points u_helmholtz_list = jax.vmap(lambda xy: scheme.variables.evaluate(xy))(xy_eval) u_helmholtz = u_helmholtz_list.squeeze(-1) error_helmholtz = jnp.linalg.norm(u_helmholtz - u_exact_eval) / jnp.linalg.norm( u_exact_eval ) print(f"\nHelmholtz Equation (2D: -∇²u + k²u = f, k={k}):") print(f" Relative L2 error: {error_helmholtz:.4e}") print(f" Max pointwise error: {jnp.max(jnp.abs(u_helmholtz - u_exact_eval)):.4e}") print(f" Computed range: [{jnp.min(u_helmholtz):.4f}, {jnp.max(u_helmholtz):.4f}]") print(f" Exact range: [{jnp.min(u_exact_eval):.4f}, {jnp.max(u_exact_eval):.4f}]") # ============================================================================ # Plotting # ============================================================================ fig, axes = plt.subplots(2, 2, figsize=(12, 10)) # Plot 1: Exact solution ax = axes[0, 0] cf = ax.contourf(xx, yy, u_exact_eval.reshape(xx.shape), levels=30, cmap="turbo") plt.colorbar(cf, ax=ax) ax.set_title("Exact Solution") ax.set_xlabel("x") ax.set_ylabel("y") # Plot 2: Computed solution ax = axes[0, 1] cf = ax.contourf(xx, yy, u_helmholtz.reshape(xx.shape), levels=30, cmap="turbo") plt.colorbar(cf, ax=ax) ax.set_title("Computed Solution") ax.set_xlabel("x") ax.set_ylabel("y") # Plot 3: Pointwise error ax = axes[1, 0] error_map = u_helmholtz - u_exact_eval cf = ax.contourf(xx, yy, error_map.reshape(xx.shape), levels=30, cmap="turbo") plt.colorbar(cf, ax=ax) ax.set_title(f"Pointwise Error (L2: {error_helmholtz:.2e})") ax.set_xlabel("x") ax.set_ylabel("y") # Plot 4: Collocation points ax = axes[1, 1] ax.scatter(xy_centers[:, 0], xy_centers[:, 1], s=5, alpha=0.5, label="Interior") ax.scatter( xy_centers_bc[:, 0], xy_centers_bc[:, 1], s=5, alpha=0.5, color="red", label="Boundary", ) ax.set_title("Collocation Points") ax.set_xlabel("x") ax.set_ylabel("y") ax.legend() ax.set_aspect("equal") plt.tight_layout() plt.savefig("/tmp/helmholtz_2d_solution.png", dpi=150) print("\nPlot saved to /tmp/helmholtz_2d_solution.png")