r"""Solves an n-D advection-diffusion equation using Neural Semi-Lagrangian (NSL). .. math:: \partial_t u + a \cdot \nabla u - \sigma \Delta u & = 0 \quad \text{in } \Omega \times (0, T) \\ u & = u_0 \quad \text{on } \Omega \times \{0\} \\ u & \text{ periodic on } \partial \Omega where :math:`u: \Omega \times (0, T) \to \mathbb{R}` is the unknown function, :math:`\Omega = (-1, 1)^d \subset \mathbb{R}^d` is the spatial domain, :math:`(0, T) \subset \mathbb{R}` is the time domain, :math:`a` is a constant advection velocity and :math:`\sigma` is a constant, isotropic diffusion coefficient. The exact solution, advection speed, diffusion coefficient and frequency are ported unchanged from the torch version of this example (`examples/examples_torch/time_discrete/advection_diffusion_nd.py`), so results are comparable across the two backends. The spatial dimension `DIM` is a free parameter: the same scheme runs unmodified from 1D to n-D. This example demonstrates: - Constant, isotropic diffusion (`NeuralSemiLagrangian`'s `diffusion_coefficient_x` argument) on top of advection. - Periodic boundary conditions on an n-D `HypercubeND` domain. - Tracking errors over time against a known exact solution. See `advection_diffusion_1d_1v.py` for diffusion in the "v" component (and combined x/v diffusion) on a kinetic-style phase-space problem. """ import math import time import jax import jax.numpy as jnp import matplotlib.pyplot as plt from scimba_jax.domains.meshless_domains.domains_nd import HypercubeND from scimba_jax.nonlinear_approximation.approximation_spaces.approximation_spaces import ( # noqa: E501 ApproximationSpace, ) from scimba_jax.nonlinear_approximation.integration.monte_carlo import ( DomainSampler, TensorizedSampler, ) from scimba_jax.nonlinear_approximation.model_class.funcparam_vectorial import ( ParamVecFunction, ) from scimba_jax.nonlinear_approximation.networks.mlp import MLP from scimba_jax.nonlinear_approximation.numerical_solvers.neural_semi_lagrangian import ( NeuralSemiLagrangian, ) DIM = 2 # spatial dimension; the scheme below is unchanged for any DIM >= 1 N_COLLOC = 2000 * DIM N_EPOCHS_INIT = 300 N_EPOCHS = 30 NT = 16 # number of time steps DIFFUSION_COEFFICIENT = 0.1 VELOCITY = 1.0 # advection speed, identical in every direction FREQUENCY = jnp.pi DOMAIN_BOUNDS = ((-1.0,), (1.0,)) # isotropic: same bounds along every axis DOM_X = HypercubeND([(-1.0, 1.0)] * DIM, is_main_domain=True) SAMPLER = TensorizedSampler([DomainSampler(DOM_X)], model_type="x") # End time set so that the diffusive decay halves the amplitude of the # solution's oscillating part, as in the torch example. END_TIME = math.log(2.0) / (DIM * DIFFUSION_COEFFICIENT * FREQUENCY**2) DOM_T = (0.0, END_TIME) DT = (DOM_T[1] - DOM_T[0]) / NT # Phase shift per axis, so that no two axes carry the same phase: # s_i = i / DIM, for i in 0..DIM-1. PHASE_SHIFT = jnp.linspace(0, 1 - 1 / DIM, DIM) def exact_sol(t: jnp.ndarray, x: jnp.ndarray) -> jnp.ndarray: r"""Exact solution: a single advecting/diffusing Fourier mode. .. math:: u(t, x) = 2 + \sin\Big(f \sum_i (x_i - s_i - a t)\Big) \exp(-d\, \sigma\, f^2\, t) which solves :math:`\partial_t u + a \sum_i \partial_{x_i} u - \sigma \Delta u = 0` with periodic BC on :math:`(-1, 1)^d`, where :math:`d` = DIM, :math:`a` = VELOCITY, :math:`\sigma` = DIFFUSION_COEFFICIENT, :math:`f` = FREQUENCY and :math:`s_i` = PHASE_SHIFT. Note the `- a t` term is *inside* the sum over `i` (as in the torch reference this is ported from): the phase advects at rate `DIM * VELOCITY`, not `VELOCITY`, because each of the `DIM` terms in the sum contributes its own `-a t`. Pulling `- a t` out of the sum (i.e. using `sum_i(x_i - s_i) - a t`) looks equivalent but is off by a factor of `DIM`, and desynchronizes the advected solution from the "exact" one used for error reporting and plotting. """ # x and t may arrive batched ((n, DIM), (n, 1)) or, under vmap (e.g. from # the residual/loss machinery), as single points ((DIM,), (1,)). x = jnp.atleast_2d(x) t = jnp.atleast_1d(t).reshape(x.shape[0], 1)[:, 0] arg = jnp.sum(x - PHASE_SHIFT[None, :] - VELOCITY * t[:, None], axis=1) decay = jnp.exp(-DIM * DIFFUSION_COEFFICIENT * FREQUENCY**2 * t) result = 2 + jnp.sin(FREQUENCY * arg) * decay return result if result.shape[0] == 1 else result.reshape(-1, 1) def f_init(x: jnp.ndarray) -> jnp.ndarray: """Initial condition function.""" t = jnp.zeros((1,)) if x.ndim == 1 else jnp.zeros((x.shape[0], 1)) return exact_sol(t, x) def advection_field(t: jnp.ndarray, x: jnp.ndarray) -> jnp.ndarray: """Advection field a(t, x) = VELOCITY, identical in every direction.""" return jnp.ones_like(x) * VELOCITY def exact_characteristic_foot(t: jnp.ndarray, x: jnp.ndarray, dt: float) -> jnp.ndarray: """Exact characteristic foot for constant advection: X(t) = x - a * dt.""" return x - VELOCITY * dt if __name__ == "__main__": key = jax.random.PRNGKey(0) out_size = 1 hidden_sizes = [7 * DIM] * 3 nn = MLP(in_size=DIM, out_size=out_size, hidden_sizes=hidden_sizes, key=key) space = ApproximationSpace({"x": DIM}, [(nn, "scalar", None)], model_type="x") neural_sl = NeuralSemiLagrangian( main_domain=DOM_X, time_domain=DOM_T, sampler=SAMPLER, dt=DT, advection_field=advection_field, exact_characteristic_foot=exact_characteristic_foot, periodic=True, domain_bounds=DOMAIN_BOUNDS, out_size=out_size, model_type="x", exact_solution=exact_sol, diffusion_coefficient_x=DIFFUSION_COEFFICIENT, ) print(f"Initializing the Neural Semi-Lagrangian solver in dimension {DIM}...") start_init = time.perf_counter() key, neural_sl = neural_sl.initialize( key, space, f_init, N_EPOCHS_INIT, N_COLLOC, file_name="nsl_ad_nd_init", retrain=False, ) space = neural_sl.space end_init = time.perf_counter() print(f"... done in {end_init - start_init:.2f} seconds\n") print("Solving with Neural Semi-Lagrangian method...") start_solve = time.perf_counter() key, neural_sl = neural_sl.solve(key, space, N_EPOCHS, N_COLLOC) space = neural_sl.space end_solve = time.perf_counter() print(f"... done in {end_solve - start_solve:.2f} seconds\n") errors = neural_sl.errors_over_time print(f"Initial relative L2 error: {errors[0, 0]:.2e}") print(f"Final relative L2 error: {errors[-1, 0]:.2e}") print(f"Final relative Linf error: {errors[-1, 1]:.2e}") # Visualize a 2D slice (x0, x1) of the solution on a regular grid (needed # for contourf, unlike the scattered collocation points); when DIM > 2, # the remaining coordinates are fixed to a constant, as in the torch # example. n_grid = 200 grid_1d = jnp.linspace(-1.0, 1.0, n_grid) x1, x2 = jnp.meshgrid(grid_1d, grid_1d) x_plot = jnp.stack([x1.ravel(), x2.ravel()], axis=1) if DIM > 2: slice_value = jnp.full((x_plot.shape[0], DIM - 2), 1.0 / (1.0 + DIM)) x_plot = jnp.concatenate([x_plot, slice_value], axis=1) variables = space.create_variables() all_together = ParamVecFunction.cat(variables) batched_func = all_together.vmap_on_physical_variables() u_pred = jax.device_get(batched_func(space, x_plot)).reshape(n_grid, n_grid) t_final = jnp.full((x_plot.shape[0], 1), DOM_T[-1]) u_exact = jax.device_get(exact_sol(t_final, x_plot)).reshape(n_grid, n_grid) u_ini = jax.device_get(f_init(x_plot)).reshape(n_grid, n_grid) fig, ax = plt.subplots(2, 2, figsize=(10, 10)) cf = ax[0, 0].contourf(x1, x2, u_ini, cmap="turbo", levels=100) plt.colorbar(cf, ax=ax[0, 0]) ax[0, 0].set_title("Initial condition") cf = ax[0, 1].contourf(x1, x2, u_exact, cmap="turbo", levels=100) plt.colorbar(cf, ax=ax[0, 1]) ax[0, 1].set_title(f"Exact solution at t={DOM_T[-1]:.3f}") cf = ax[1, 0].contourf(x1, x2, u_pred, cmap="turbo", levels=100) plt.colorbar(cf, ax=ax[1, 0]) ax[1, 0].set_title("Approximate solution") error = jnp.abs(u_pred - u_exact) / jnp.abs(u_exact) cf = ax[1, 1].contourf(x1, x2, error, cmap="gist_heat", levels=100) plt.colorbar(cf, ax=ax[1, 1]) ax[1, 1].set_title("Relative error") fig.suptitle(f"{DIM}D advection-diffusion, x0-x1 slice") fig.tight_layout() plt.show()