r"""Solves the Vlasov equation on a periodic square using Neural Semi-Lagrangian. .. math:: \partial_t u + v \partial_x u + \sin(x) \partial_v u = 0 where :math:`u: \mathbb{R} \times \mathbb{R} \times (0, T) \to \mathbb{R}` is the unknown function, depending on space, velocity, and time. The equation is solved using the neural semi-Lagrangian scheme with Natural Gradient preconditioning. """ # %% import os 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, ) from scimba_jax.nonlinear_approximation.integration.monte_carlo_parameters import ( UniformVelocitySamplerOnCuboid, ) 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, ) FIG_DIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), "fig") os.makedirs(FIG_DIR, exist_ok=True) # Constants for collocation points and epochs N_COLLOC = 64**2 # collocation points for 2D phase space N_EPOCHS_INIT = 250 N_EPOCHS = 250 # Domain definitions X_MIN, X_MAX = 0.0, 2 * jnp.pi V_MIN, V_MAX = -6.0, 6.0 DOMAIN_X = Segment1D((X_MIN, X_MAX), is_main_domain=True) DOMAIN_V = Segment1D((V_MIN, V_MAX)) # Domain bounds for periodic boundary conditions DOMAIN_BOUNDS = ( (X_MIN, V_MIN), # lower bounds for (x, v) (X_MAX, V_MAX), # upper bounds for (x, v) ) # Sampler for phase space (x, v) SAMPLER = TensorizedSampler( [ DomainSampler(DOMAIN_X), UniformVelocitySamplerOnCuboid(DOMAIN_V), ], model_type="x_v", ) # Time domain DOM_T = (0.0, 4.5) NT = 4 # Number of time steps DT = (DOM_T[1] - DOM_T[0]) / NT # Time step size # Natural gradient preconditioning parameters LINESEARCH = "strong-wolfe" MATRIX_REGULARIZATION = 1e-2 ADAPTIVE_MATRIX_REGULARIZATION = True ENG_KWARGS = { "linesearch": LINESEARCH, "matrix_regularization": MATRIX_REGULARIZATION, "adaptive_matrix_regularization": ADAPTIVE_MATRIX_REGULARIZATION, } def f_init(x: jnp.ndarray, v: jnp.ndarray) -> jnp.ndarray: """Initial condition function. Args: x: Spatial points of shape (1,). v: Velocity points of shape (1,). Returns: Initial values of shape (1,). """ # Initial condition: Gaussian in velocity space return jnp.exp(-(v**2) / 2) / jnp.sqrt(2 * jnp.pi) def advection_field( t: jnp.ndarray, x: jnp.ndarray, v: jnp.ndarray ) -> tuple[jnp.ndarray, jnp.ndarray]: """Advection field for Vlasov: (v, sin(x)). For the Vlasov equation: df/dt + v*df/dx + sin(x)*df/dv = 0 The characteristic ODE is: dx/dt = v dv/dt = sin(x) So the advection field returns a tuple with "dx/dt" and "dv/dt". Args: t: Time array (not used but included for signature compatibility). x: Spatial points of shape (1,). v: Velocity points of shape (1,). Returns: Tuple of arrays (dx/dt, dv/dt) representing the advection field. Each array has shape (1,). """ return v, jnp.sin(x) def exact_solution(x: jnp.ndarray, v: jnp.ndarray, t: jnp.ndarray, n_steps: int = 400): """Semi-Lagrangian reference solution (dx/dt = v, dv/dt = sin(x))""" def sl_characteristic_rhs(state): x, v = state return jnp.array(advection_field(None, x, v)) def sl_rk4_step(state, dt): k1 = sl_characteristic_rhs(state) k2 = sl_characteristic_rhs(state + 0.5 * dt * k1) k3 = sl_characteristic_rhs(state + 0.5 * dt * k2) k4 = sl_characteristic_rhs(state + dt * k3) return state + (dt / 6.0) * (k1 + 2.0 * k2 + 2.0 * k3 + k4) def sl_backward_flow(x, v, t, n_steps): dt = -t / n_steps def body(state, _): return sl_rk4_step(state, dt), None state, _ = jax.lax.scan(body, jnp.array([x, v]), None, length=n_steps) return state def sl_density(x, v, t, n_steps): _, v0 = sl_backward_flow(x, v, t, n_steps) return f_init(jnp.array([0.0]), jnp.array([v0])) u_func = jax.jit( jax.vmap( jax.vmap(sl_density, in_axes=(0, 0, None, None)), in_axes=(0, 0, None, None) ), static_argnums=(3,), ) return u_func(x, v, t, n_steps) def plot_solution( nsl: NeuralSemiLagrangian, space: ApproximationSpace, t: float, key: jax.random.PRNGKey, n_points: int = 512, ) -> None: """Plot the numerical solution using contourf at a given time. Args: nsl: The NeuralSemiLagrangian solver instance. space: The current approximation space. t: Time at which to plot the solution. key: Random key for sampling (not used, kept for compatibility). n_points: Number of points to use for plotting (default: 512 per dimension). """ # Create a regular meshgrid for plotting x_min, x_max = X_MIN, X_MAX v_min, v_max = V_MIN, V_MAX x_vals = jnp.linspace(x_min, x_max, n_points) v_vals = jnp.linspace(v_min, v_max, n_points) X_grid, V_grid = jnp.meshgrid(x_vals, v_vals, indexing="ij") # Evaluate solution at grid points # Reshape to (n_points, 1) for each component x_flat = X_grid.flatten()[:, jnp.newaxis] v_flat = V_grid.flatten()[:, jnp.newaxis] # Define plotting function def plot_u(ax, u, t_val, kind, relative_l2=None): cmap = "inferno" if kind == "|NSL - SL|" else "turbo" im = ax.contourf(x_vals, v_vals, u.T, levels=256, cmap=cmap, zorder=-20) plt.colorbar(im, ax=ax) ax.contour( im, levels=im.levels[:: 256 // 6], colors="w", alpha=0.5, linewidths=0.8, zorder=-20, ) ax.set_rasterization_zorder(-10) if kind in ["NSL", "SL"]: title = f"{kind} density t={t_val} (min={float(u.min()):.2f}, max={float(u.max()):.2f})" else: title = rf"{kind}: t={t_val} (rel. $L^2$={float(relative_l2):.2e})" ax.set_title(title) # Build input args for the model args = (x_flat, v_flat) # Evaluate predicted solution 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, *args)) # Reshape back to 2D grid u_grid = u_pred.reshape(n_points, n_points) # Evaluate exact solution u_ref = exact_solution(x_flat, v_flat, t).reshape(n_points, n_points) # Evaluate error u_error = jnp.abs(u_grid - u_ref) relative_l2 = jnp.sqrt(jnp.mean(u_error**2)) / jnp.sqrt(jnp.mean(u_ref**2)) # Plot the three quantities for qty, kind in zip([u_grid, u_ref, u_error], ["NSL", "SL", "|NSL - SL|"]): fig, ax = plt.subplots(1, 1, figsize=(5, 4)) plot_u(ax, qty, t, kind, relative_l2) plt.tight_layout() plt.savefig(os.path.join(FIG_DIR, f"{kind}_density_pinn_last_time.pdf")) plt.show() # %% if __name__ == "__main__": """Solve the Vlasov equation using Neural Semi-Lagrangian.""" key = jax.random.PRNGKey(0) # Build the Neural Semi-Lagrangian solver nsl = NeuralSemiLagrangian( main_domain=DOMAIN_X, time_domain=DOM_T, sampler=SAMPLER, dt=DT, advection_field=advection_field, periodic=True, domain_bounds=DOMAIN_BOUNDS, out_size=1, model_type="x_v", n_rk_steps=10, ) # Create a simple MLP for the approximation space # For model_type="x_v", input size is dim_x + dim_v = 1 + 1 = 2 key, subkey = jax.random.split(key) nn = MLP(in_size=2, out_size=1, hidden_sizes=[16] * 4, key=subkey) space = ApproximationSpace( {"x": 1, "v": 1}, [(nn, "scalar", None)], model_type="x_v" ) # Initialize the Neural Semi-Lagrangian solver print("Initializing the Neural Semi-Lagrangian solver...") start_init = time.perf_counter() key, nsl = nsl.initialize( key, space, f_init, N_EPOCHS_INIT, N_COLLOC, file_name="kinetic_1d_1v_vlasov", retrain=False, **ENG_KWARGS, ) space = nsl.space end_init = time.perf_counter() print( f"Initializing the Neural Semi-Lagrangian solver... " f"Done in {end_init - start_init:.2f} seconds\n" ) # Solve with Neural Semi-Lagrangian method print("Solving with Neural Semi-Lagrangian method...") start_solve = time.perf_counter() key, nsl = nsl.solve(key, space, N_EPOCHS, N_COLLOC, **ENG_KWARGS) space = nsl.space end_solve = time.perf_counter() print( f"Solving with Neural Semi-Lagrangian method... " f"Done in {end_solve - start_solve:.2f} seconds\n" ) # Print final solution info print("Solution computed successfully!") print(f"Number of time steps: {nsl.nt}") print(f"Errors over time shape: {nsl.errors_over_time.shape}") # Print final errors if available if hasattr(nsl, "errors_over_time"): errors = nsl.errors_over_time print(f"\nFinal relative L2 error: {errors[-1, 0]:.2e}") print(f"Final relative Linf error: {errors[-1, 1]:.2e}") # Plot the solution at final time print("\nPlotting solution at final time...") plot_solution(nsl, space, DOM_T[1], key) plt.show() # %%