"""Solve a 2D advection-diffusion-reaction (ADR) equation in a convection-dominated regime using kernel-based collocation. The PDE is given by: -∇·(d∇u) + b·∇u + c*u = f in Ω = [-1,1]² u = g on ∂Ω Manufactured solution: u(x,y) = sin(π*x)*sin(π*y) with coefficients: d = I (isotropic diffusion) b = (10, 0) (strong advection in x direction - convection-dominated) c = 1 (reaction) g = 0 (homogeneous Dirichlet BC) RHS f is computed from the PDE: f = -∇·(d∇u) + b·∇u + c*u High Péclet number (Pe = ||b||*L/d ≈ 10) makes this a convection-dominated problem, which tests solver stability with large advection terms. """ 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.general_elliptic import GeneralElliptic # N_COLLOC = 1_500 # More collocation points for convection-dominated regime N_COLLOC = 500 # 1_500 points id too much N_BC_COLLOC = 300 key = jax.random.PRNGKey(0) # Péclet number ~ 10 (convection-dominated) ADVECTION_STRENGTH = 10.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 b_fn(x: jnp.ndarray, mu=None) -> jnp.ndarray: """Advection velocity: b = (ADVECTION_STRENGTH, 0) - strong advection in x direction""" return jnp.array([ADVECTION_STRENGTH, 0.0]) def c_fn(x: jnp.ndarray, mu=None) -> jnp.ndarray: """Reaction coefficient: c = 1""" return 1.0 def f_rhs(x: jnp.ndarray, mu=None) -> jnp.ndarray: """ RHS for: -∇·(d∇u) + b·∇u + c*u = f With u = sin(π*x)*sin(π*y) and b = (ADVECTION_STRENGTH, 0): ∂u/∂x = π*cos(π*x)*sin(π*y) ∂u/∂y = π*sin(π*x)*cos(π*y) ∂²u/∂x² = -π²*sin(π*x)*sin(π*y) ∂²u/∂y² = -π²*sin(π*x)*sin(π*y) ∇²u = -2π²*sin(π*x)*sin(π*y) d(x,y) = I f = -∇²u + b·∇u + c*u = 2π²*sin(π*x)*sin(π*y) + ADVECTION_STRENGTH*π*cos(π*x)*sin(π*y) + 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) du_dx = jnp.pi * jnp.cos(jnp.pi * x_coord) * jnp.sin(jnp.pi * y_coord) du_dy = jnp.pi * jnp.sin(jnp.pi * x_coord) * jnp.cos(jnp.pi * y_coord) # noqa F841 d2u_dx2 = -(jnp.pi**2) * jnp.sin(jnp.pi * x_coord) * jnp.sin(jnp.pi * y_coord) d2u_dy2 = -(jnp.pi**2) * jnp.sin(jnp.pi * x_coord) * jnp.sin(jnp.pi * y_coord) laplacian_u = d2u_dx2 + d2u_dy2 f = -laplacian_u + ADVECTION_STRENGTH * du_dx + 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 ADR convection-dominated, Pe≈{ADVECTION_STRENGTH})..." ) # 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) du_dx_exact = jnp.pi * jnp.cos(jnp.pi * x) * jnp.sin(jnp.pi * y) du_dy_exact = jnp.pi * jnp.sin(jnp.pi * x) * jnp.cos(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 + ADVECTION_STRENGTH*∂u/∂x + u = f residual_val = -laplacian_exact + ADVECTION_STRENGTH * du_dx_exact + 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] xy_normals_bc = sample_dict["boundary"][1] 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.4 ), # Smaller sigma for better stability in convection-dominated regime centers=xy_centers, learnable_center_bool=False, basis_type="scalar", ) variables = CollocationVariables( basis=basis, nb_variables=1, ) # ============================================================================ # Create ADR PDE model # ============================================================================ # # -∇·(d∇u) + b·∇u + c*u = f with weak Dirichlet BC is exactly what # GeneralElliptic already implements (alpha=0, beta=1 reduces its Robin BC to # Dirichlet) -- no need to redefine it here. Weak Dirichlet needs a real # boundary normal even though it never uses one mathematically (RobinResidual # always reads it, alpha=0 just zeroes its contribution), so it is passed # through as `bc_normals` below -- TensorizedSampler(bc=True) already # computes it, it was simply unused before. pde = GeneralElliptic( main_domain=dx, model_type="x_mu", f_rhs=f_rhs, bc="weak", f_bc_rhs=f_bc, A=d_fn, b=b_fn, c=c_fn, alpha=0.0, beta=1.0, ) scheme = EllipticCollocationScheme( pde=pde, variables=variables, collocation_points=xy_centers, bc_collocation_points=xy_centers_bc, bc_normals=xy_normals_bc, ) # ============================================================================ # Solve # ============================================================================ scheme.variables.use_scan = True print(f"\nAssembling system (ADR 2D convection-dominated, Pe≈{ADVECTION_STRENGTH})...") A_mat = scheme.assembly_scheme() print(f"Assembly matrix shape: {A_mat.shape}") print("Solving system (ADR 2D convection-dominated)...") 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_adr_list = jax.vmap(lambda xy: scheme.variables.evaluate(xy))(xy_eval) u_adr = u_adr_list.squeeze(-1) error_adr = jnp.linalg.norm(u_adr - u_exact_eval) / jnp.linalg.norm(u_exact_eval) print( f"\nADR Convection-Dominated (2D: -∇·(d∇u) + b·∇u + c*u = f, b_x={ADVECTION_STRENGTH}):" ) print(f" Relative L2 error: {error_adr:.4e}") print(f" Max pointwise error: {jnp.max(jnp.abs(u_adr - u_exact_eval)):.4e}") print(f" Computed range: [{jnp.min(u_adr):.4f}, {jnp.max(u_adr):.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_adr.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_adr - 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_adr:.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/adr_2d_convection_dominated.png", dpi=150) print("\nPlot saved to /tmp/adr_2d_convection_dominated.png")