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 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 N_COLLOC = 1000 N_BC_COLLOC = 100 key = jax.random.PRNGKey(0) def f_rhs(xy: jnp.ndarray, mu) -> jnp.ndarray: x, y = xy[0:1], xy[1:2] return 2 * jnp.pi**2 * jnp.sin(jnp.pi * x) * jnp.sin(jnp.pi * y) def f_bc(xy: jnp.ndarray, n, mu) -> jnp.ndarray: x, _ = xy[0:1], xy[1:2] return jnp.zeros_like(x) def exact_sol(xy: jnp.ndarray, mu) -> jnp.ndarray: x, y = xy[:, 0:1], xy[:, 1:2] return jnp.sin(jnp.pi * x) * jnp.sin(jnp.pi * y) domain_x = [(-1.0, 1.0), (-1.0, 1.0)] dx = Square2D(domain_x, is_main_domain=True) sampler = TensorizedSampler([DomainSampler(dx)], bc=True) # sample the domain key, sample_dict = sampler.sample(key, N_COLLOC, N_BC_COLLOC) pde = LaplacianDirichletND( dx, lambda *args: f_rhs(*args), bc="weak", f_bc_rhs=lambda *args: f_bc(*args) ) x_centers = sample_dict["interior"][0] x_centers_bc = sample_dict["boundary"][0] basis = KernelBasis( dim=2, output_dim=1, kernel_function=GaussianKernel(sigma=0.5), centers=x_centers, learnable_center_bool=False, basis_type="scalar", ) variables = CollocationVariables( basis=basis, nb_variables=1, ) scheme = EllipticCollocationScheme( pde=pde, variables=variables, collocation_points=x_centers, bc_collocation_points=x_centers_bc, ) # print(scheme.pde.physical_residuals["boundary"].construct_residual) A = scheme.assembly_scheme() print(A.shape) scheme = EllipticCollocationScheme.solve(scheme) # Evaluate the solution # create test grid for evaluation x_eval = jnp.linspace(-1, 1, 50) y_eval = jnp.linspace(-1, 1, 50) xx, yy = jnp.meshgrid(x_eval, y_eval) xy_eval = jnp.stack([xx.flatten(), yy.flatten()], axis=-1) var = scheme.variables u_kernel = jax.vmap(lambda x: var.evaluate(x))(xy_eval) # True solution for the PDE with f=1 and Dirichlet BCs u_true = exact_sol(xy_eval, None) error = jnp.linalg.norm(u_kernel - u_true) / jnp.linalg.norm(u_true) print(f"Relative L2 error at collocation points: {error:.2e}") # Plot the solution side by side with the true solution plt.figure(figsize=(12, 5)) plt.subplot(1, 2, 1) plt.contourf(xx, yy, u_kernel.reshape(xx.shape), levels=50, cmap="turbo") plt.colorbar() plt.title("Kernel solution") plt.subplot(1, 2, 2) plt.contourf(xx, yy, u_true.reshape(xx.shape), levels=50, cmap="turbo") plt.colorbar() plt.title("True solution") plt.show()