"""Solve a 1D advection-diffusion-reaction (ADR) equation using kernel-based collocation. The PDE is given by: -(d(x)*u')' + b*u' + c*u = f in Ω = [-1,1] u = g on ∂Ω Manufactured solution: u(x) = sin(π*x) with coefficients: d(x) = 1 + 0.5*x (diffusion) b = 1 (advection) c = 1 (reaction) RHS f is computed from the PDE: f = -(d(x)*u')' + b*u' + c*u = -(d'(x)*u' + d(x)*u'') + b*u' + c*u This is -∇·(d∇u) + b·∇u + c*u in divergence form -- the same operator every other ADR example in this directory already gets from ``GeneralEllipticResidual``/``GeneralElliptic`` (``scimba_jax.physical_models.elliptic_pde.general_elliptic``), reused here too instead of a hand-rolled residual. An earlier version of this example used the non-divergence form ``-d(x)*u'' + b*u' + c*u`` (missing the ``-d'(x)*u'`` term), a genuinely different PDE that ``GeneralEllipticResidual`` does not implement; the manufactured solution's RHS ``f`` above is the one that matches the divergence form instead, since a non-constant ``d(x)`` makes the two forms disagree. """ import jax import jax.numpy as jnp import matplotlib.pyplot as plt from scimba_jax.domains.meshless_domains.domains_1d import Segment1D 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 = 500 N_BC_COLLOC = 2 key = jax.random.PRNGKey(0) # ============================================================================ # Manufactured Solution and Problem Definition # ============================================================================ def u_exact(x: jnp.ndarray, mu=None) -> jnp.ndarray: """Exact solution: u(x) = sin(π*x)""" if x.ndim == 0: x = x[None] return jnp.sin(jnp.pi * x).squeeze() def d_fn(x: jnp.ndarray, mu=None) -> jnp.ndarray: """Diffusion coefficient: d(x) = 1 + 0.5*x, as the 1x1 matrix anisotropic_laplacian expects.""" x_val = x if x.ndim == 0 else x[0] return (1.0 + 0.5 * x_val).reshape(1, 1) def b_fn(x: jnp.ndarray, mu=None) -> jnp.ndarray: """Advection velocity: b = (1,)""" return jnp.array([1.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(x)*u')' + b*u' + c*u = f With u = sin(π*x): u' = π*cos(π*x) u'' = -π²*sin(π*x) d(x) = 1 + 0.5*x, d'(x) = 0.5 -(d*u')' = -(d'*u' + d*u'') = -0.5*π*cos(π*x) + (1+0.5*x)*π²*sin(π*x) f = -(d*u')' + b*u' + c*u = -0.5*π*cos(πx) + (1+0.5x)*π²*sin(πx) + π*cos(πx) + sin(πx) """ x_val = x if x.ndim == 0 else x[0:1] u = jnp.sin(jnp.pi * x_val) du_dx = jnp.pi * jnp.cos(jnp.pi * x_val) d2u_dx2 = -(jnp.pi**2) * jnp.sin(jnp.pi * x_val) d = 1.0 + 0.5 * x_val d_prime = 0.5 diffusion = -(d_prime * du_dx + d * d2u_dx2) f = diffusion + du_dx + u # b=1, c=1 return f def f_bc(x: jnp.ndarray, n, mu=None) -> jnp.ndarray: """Boundary condition: u = 0 (homogeneous Dirichlet)""" # Handle both scalar (for verification) and array (for vmapped assembly) if x.ndim == 0: return jnp.zeros_like(x) else: x_val = x[0:1] # Keep dimension like in kernel_example return jnp.zeros_like(x_val) # Return shape (1,) # ============================================================================ # Verify Manufactured Solution # ============================================================================ print("Verifying manufactured solution (1D ADR with diffusion)...") # Test at a single point x_test = jnp.array(0.5) # Compute exact solution and its derivatives u_val = u_exact(x_test) print(f"u(0.5) = {u_val:.6f}") # Compute derivatives manually x = 0.5 u_exact_val = jnp.sin(jnp.pi * x) du_dx_exact = jnp.pi * jnp.cos(jnp.pi * x) d2u_dx2_exact = -(jnp.pi**2) * jnp.sin(jnp.pi * x) d_coeff = 1.0 + 0.5 * x d_prime_coeff = 0.5 # Verify PDE: -(d(x)*u')' + b*u' + c*u = f residual_val = ( -(d_prime_coeff * du_dx_exact + d_coeff * d2u_dx2_exact) + du_dx_exact + u_exact_val ) f_val = f_rhs(x_test) print(f"d(0.5) = {d_coeff:.6f}") print(f"f(0.5) = {f_val:.6f}") print(f"-(d(x)*u')' + b*u' + c*u = {residual_val:.6f}") print(f"Residual error (should be 0): {jnp.abs(residual_val - f_val):.6e}") print() # ============================================================================ domain_x = jnp.array([[-1.0, 1.0]]) dx = Segment1D(domain_x, is_main_domain=True) sampler = TensorizedSampler([DomainSampler(dx)], bc=True) key, sample_dict = sampler.sample(key, N_COLLOC, N_BC_COLLOC) x_centers = sample_dict["interior"][0] x_centers_bc = sample_dict["boundary"][0] x_normals_bc = sample_dict["boundary"][1] # ============================================================================ # Setup Basis and Variables # ============================================================================ basis = KernelBasis( dim=1, output_dim=1, kernel_function=GaussianKernel(sigma=0.2), centers=x_centers, learnable_center_bool=False, basis_type="scalar", ) variables = CollocationVariables( basis=basis, nb_variables=1, ) # ============================================================================ # Create ADR PDE model # ============================================================================ # # -(d(x)*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=x_centers, bc_collocation_points=x_centers_bc, bc_normals=x_normals_bc, ) # ============================================================================ # Solve # ============================================================================ print("Assembling system (GeneralEllipticResidual)...") A_mat = scheme.assembly_scheme() print(f"Assembly matrix shape: {A_mat.shape}") print("Solving system (GeneralEllipticResidual)...") scheme = EllipticCollocationScheme.solve(scheme) print("Done!") # ============================================================================ # Evaluate and Compare # ============================================================================ n_eval = 200 # 1D evaluation points - easy with scan! x_eval_1d = jnp.linspace(-1, 1, n_eval) # Reshape to (n_eval, 1) to match 1D domain format x_eval = x_eval_1d.reshape(-1, 1) print("Evaluating solutions...") u_exact_eval = jax.vmap(u_exact)(x_eval_1d) # Use scan for evaluation (memory efficient) - vmap over the 1D points u_adr_list = jax.vmap(lambda x: scheme.variables.evaluate(x))(x_eval) u_adr = u_adr_list.squeeze(-1) # Remove output dimension if needed error_adr = jnp.linalg.norm(u_adr - u_exact_eval) / jnp.linalg.norm(u_exact_eval) print("\nGeneralEllipticResidual (1D: -(d(x)*u')' + b*u' + c*u = f):") 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(1, 2, figsize=(12, 4)) # Plot 1: Exact vs Computed ax = axes[0] ax.plot(x_eval, u_exact_eval, "b-", linewidth=2, label="Exact") ax.plot(x_eval, u_adr, "r--", linewidth=2, label="Computed") ax.set_xlabel("x") ax.set_ylabel("u") ax.set_title("1D ADR Solution with Variable Diffusion") ax.legend() ax.grid(True) # Plot 2: Pointwise error ax = axes[1] error = u_adr - u_exact_eval ax.plot(x_eval, error, "g-", linewidth=2) ax.set_xlabel("x") ax.set_ylabel("error") ax.set_title(f"Pointwise Error (L2 error: {error_adr:.2e})") ax.grid(True) plt.tight_layout() plt.savefig("/tmp/adr_1d_solution_diffusion.png", dpi=150) print("\nPlot saved to /tmp/adr_1d_solution_diffusion.png")