r"""Solves the advection equation in 1D using Neural Semi-Lagrangian (NSL) with a parameter (advection speed). .. math:: \partial_t u + c \partial_{x} u & = 0 \quad \text{in } \Omega \times (0, T) \\ u & = u_0 \quad \text{on } \Omega \times \{0\} \\ u(0, t) & = u(1, t) \quad \text{periodic BC} where :math:`u: \Omega \times (0, T) \to \mathbb{R}` is the unknown function, :math:`\Omega \subset \mathbb{R}` is the spatial domain and :math:`(0, T) \subset \mathbb{R}` is the time domain. The equation is solved on a segment domain with periodic boundary conditions using the Neural Semi-Lagrangian method. The advection velocity is a parameter (c) that can vary. This example demonstrates: - Using periodic boundary conditions via modulo operation on characteristics - Solving the advection equation with a parametric velocity field - Comparing the NSL solution to the exact solution - Parameter sweep with different advection speeds """ import time 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.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 = 2000 N_EPOCHS_INIT = 125 N_EPOCHS = 25 DOMAIN_BOUNDS = ((0.0,), (1.0,)) X_MIN, X_MAX = DOMAIN_BOUNDS[0][0], DOMAIN_BOUNDS[1][0] DOM_X = Segment1D((X_MIN, X_MAX), is_main_domain=True) DOM_T = (0.0, 1.5) # Parameter domain: [c_min, c_max] for advection speed # For a parameter sweep, set c_min < c_max # For a fixed case, set c_min == c_max C_MIN, C_MAX = 0.3, 0.7 # Advection speed range DOM_MU = [[C_MIN, C_MAX]] NT = 15 # Number of time steps DT = (DOM_T[1] - DOM_T[0]) / NT # Time step size PARAMS_TO_PLOT = [[0.3], [0.5], [0.7]] # Advection speeds to plot (each as list) def exact_sol(t: jnp.ndarray, x: jnp.ndarray, mu: jnp.ndarray) -> jnp.ndarray: """Exact solution of the 1D advection equation with periodic BC and parameter. The exact solution is a Gaussian profile advected at velocity c: u(x, t) = exp(-150 * (x - (0.25 + c*t))^2) With periodic boundary conditions, the profile wraps around when it reaches the domain boundaries. Args: t: Time array of shape (n_points, 1). x: Spatial points of shape (n_points, 1). mu: Parameter array of shape (n_points, 1) where mu[:, 0] = c. 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 parameter (advection speed) c = mu[:, 0] # Compute the raw distance from the moving center raw_distance = x[:, 0] - (0.25 + c * t[:, 0]) # Apply periodic boundary conditions to the distance itself. L = X_MAX - X_MIN wrapped_distance = (raw_distance + L / 2) % L - L / 2 # Compute the exact solution using the wrapped distance result = jnp.exp(-(wrapped_distance**2) * 150) 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, 1). mu: Parameter array of shape (n_points, 1). 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, mu) = c (parametric velocity). Args: t: Time array. x: Spatial points of shape (n_points, 1). mu: Parameter array of shape (n_points, 1) where mu[:, 0] = c. Returns: Advection velocities of shape (n_points, 1). """ # mu[:, 0] contains the advection speed c return jnp.ones_like(x) * mu if __name__ == "__main__": # Initialize the key and neural network key = jax.random.PRNGKey(0) in_size = 2 # x, c out_size = 1 nn = MLP(in_size=in_size, out_size=out_size, hidden_sizes=[12] * 2, key=key) # Use model_type="x_mu" for parametric problems space = ApproximationSpace( {"x": 1, "mu": 1}, [(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, periodic=True, domain_bounds=DOMAIN_BOUNDS, 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) end_init = time.perf_counter() space = neural_sl.space losses = neural_sl.losses 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, loss=losses, 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) end_solve = time.perf_counter() space = neural_sl.space losses = neural_sl.losses 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), loss=losses, 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()