r"""Solves the advection equation in 2D using Neural Semi-Lagrangian (NSL). .. math:: \partial_t u + \mathbf{a} \cdot \nabla u & = 0 \quad \text{in } \Omega \times (0, T) \\ u & = u_0 \quad \text{on } \Omega \times \{0\} where :math:`u: \Omega \times (0, T) \to \mathbb{R}` is the unknown function, :math:`\Omega \subset \mathbb{R}^2` is the spatial domain and :math:`(0, T) \subset \mathbb{R}` is the time domain. The equation is solved on a disk domain with rotating velocity field using the Neural Semi-Lagrangian method. The velocity field is: :math:`\mathbf{a}(x) = (-2\pi x_2, 2\pi x_1)` which gives a rigid rotation. This example demonstrates: - Solving the advection equation on a 2D disk domain - Using a rotating velocity field - Tracking errors over time """ import time import jax import jax.numpy as jnp import matplotlib.pyplot as plt from scimba_jax.domains.meshless_domains.domains_2d import Disk2D from scimba_jax.nonlinear_approximation.approximation_spaces.approximation_spaces import ( ApproximationSpace, ) from scimba_jax.nonlinear_approximation.integration.monte_carlo import ( DomainSampler, TensorizedSampler, UniformParametricSampler, ) from scimba_jax.nonlinear_approximation.networks.mlp import MLP from scimba_jax.nonlinear_approximation.numerical_solvers.neural_semi_lagrangian import ( NeuralSemiLagrangian, ) from scimba_jax.plots.plots_nd import plot_abstract_approx_spaces N_COLLOC = 3000 N_EPOCHS_INIT = 500 N_EPOCHS = 50 # Disk domain with center (0, 0) and radius 1 DISK_CENTER = (0.0, 0.0) DISK_RADIUS = 1.0 DOM_X = Disk2D(DISK_CENTER, DISK_RADIUS, is_main_domain=True) DOM_T = (0.0, 0.5) NT = 2 # Number of time steps DT = (DOM_T[1] - DOM_T[0]) / NT # Time step size # Parameter domain: [c_min, c_max], [v_min, v_max] # For a fixed case, set min == max C_MIN, C_MAX = 0.2, 0.4 # Initial bump center V_MIN, V_MAX = 0.05, 0.1 # Variance parameter DOM_MU = [[C_MIN, C_MAX], [V_MIN, V_MAX]] PARAMS_TO_PLOT = [[0.3, 0.075], [0.22, 0.09], [0.38, 0.06]] def exact_sol(t: jnp.ndarray, x: jnp.ndarray, mu: jnp.ndarray) -> jnp.ndarray: """Exact solution of the 2D rotating advection equation with parameters. The exact solution represents a Gaussian bump rotating rigidly around the origin with angular velocity 2*pi. The initial bump is centered at (c, 0) where c is a parameter, and rotates according to the flow: x1(t) = c * cos(2*pi*t) x2(t) = c * sin(2*pi*t) For a point x at time t, we find where it was at time 0 by rotating backwards: x0 = R(-2*pi*t) * x where R(theta) is the rotation matrix. Args: t: Time array of shape (n_points, 1). x: Spatial points of shape (n_points, 2). mu: Parameters of shape (n_points, 2) where mu[:, 0] = c (center) and mu[:, 1] = v (variance). Returns: Solution values of shape (n_points, 1). """ # Ensure correct shapes x = jnp.atleast_2d(x) t = jnp.atleast_2d(t) mu = jnp.atleast_2d(mu) # Extract parameters c = mu[:, 0] v = mu[:, 1] # Compute the rotation angle (backwards in time angle = -2.0 * jnp.pi * t[:, 0] # x has shape (n_points, 2) x1, x2 = x[:, 0], x[:, 1] x1_0 = jnp.cos(angle) * x1 - jnp.sin(angle) * x2 x2_0 = jnp.sin(angle) * x1 + jnp.cos(angle) * x2 # Initial bump centered at (c, 0) with variance v dist_sq = (x1_0 - c) ** 2 + x2_0**2 result = 1.0 + jnp.exp(-dist_sq / (2.0 * v**2)) if result.shape[0] == 1: return result else: return result.reshape(-1, 1) # Return shape (n_points, 1) def f_init(x: jnp.ndarray, mu: jnp.ndarray) -> jnp.ndarray: """Initial condition function. Args: x: Spatial points of shape (n_points, 2). mu: Parameters of shape (n_points, 2). Returns: Initial values of shape (n_points, 1). """ if x.ndim == 1: t = jnp.zeros((1,)) else: t = jnp.zeros((x.shape[0], 1)) return exact_sol(t, x, mu) def advection_field(t: jnp.ndarray, x: jnp.ndarray, mu: jnp.ndarray) -> jnp.ndarray: """Advection field a(t, x) = (-2*pi*x2, 2*pi*x1) (rotating field). This gives a rigid rotation with angular velocity 2*pi. """ return jnp.stack([-2.0 * jnp.pi * x[1], 2.0 * jnp.pi * x[0]]) def exact_characteristic_foot( t: jnp.ndarray, x: jnp.ndarray, mu: jnp.ndarray, dt: float ) -> jnp.ndarray: """Exact characteristic foot computation. For the rotating advection equation u_t + a·∇u = 0 with a(x) = (-2*pi*x2, 2*pi*x1), the characteristics are: X(s) = R(2*pi*(s-t)) * x where R(theta) is the rotation matrix. We want the foot at time t given x at time t+dt: X(t) = R(-2*pi*dt) * x Args: t: Time array of shape (1,). x: Spatial points of shape (2,). mu: Parameters of shape (2,). dt: Time step size. Returns: Foot points of shape (2,). """ # Rotation angle for backwards characteristic # dt is a scalar angle = -2.0 * jnp.pi * dt # x has shape (2,) x1, x2 = x[0], x[1] x1_foot = jnp.cos(angle) * x1 - jnp.sin(angle) * x2 x2_foot = jnp.sin(angle) * x1 + jnp.cos(angle) * x2 return jnp.stack([x1_foot, x2_foot]) if __name__ == "__main__": # Initialize the key and neural network key = jax.random.PRNGKey(0) in_size = 4 # x1, x2, c, v out_size = 1 nn = MLP(in_size=in_size, out_size=out_size, hidden_sizes=[12] * 4, key=key) # Use model_type="x_mu" for parametric problems space = ApproximationSpace( {"x": 2, "mu": 2}, [(nn, "scalar", None)], model_type="x_mu" ) # Create the neural semi-Lagrangian solver # Use TensorizedSampler with both spatial and parametric samplers sampler = TensorizedSampler( [DomainSampler(DOM_X), UniformParametricSampler(DOM_MU)], model_type="x_mu", ) 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=False, domain_bounds=None, out_size=out_size, model_type="x_mu", exact_solution=exact_sol, ) print("Initializing the Neural Semi-Lagrangian solver...") start_init = time.perf_counter() key, neural_sl = neural_sl.initialize(key, space, f_init, N_EPOCHS_INIT, N_COLLOC) space = neural_sl.space end_init = time.perf_counter() print( f"Initializing the Neural Semi-Lagrangian solver... " f"Done in {end_init - start_init:.2f} seconds\n" ) # Plot initial condition plot_abstract_approx_spaces( [space], DOM_X, DOM_MU, solution=f_init, error=f_init, parameters_values=PARAMS_TO_PLOT, title="initial condition", ) 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"Solving with Neural Semi-Lagrangian method... " f"Done in {end_solve - start_solve:.2f} seconds\n" ) # Plot final solution plot_abstract_approx_spaces( [space], DOM_X, DOM_MU, solution=lambda x, mu: exact_sol(jnp.ones((x.shape[0], 1)) * DOM_T[-1], x, mu), error=lambda x, mu: exact_sol(jnp.ones((x.shape[0], 1)) * DOM_T[-1], x, mu), parameters_values=PARAMS_TO_PLOT, title=f"solution at final time t={DOM_T[-1]}", ) # Print final errors if hasattr(neural_sl, "errors_over_time"): errors = neural_sl.errors_over_time print(f"\nFinal relative L2 error: {errors[-1, 0]:.2e}") print(f"Final relative Linf error: {errors[-1, 1]:.2e}") plt.show()