r"""Shows the effect of "x"- and "v"-diffusion on a 1D-1V phase-space problem. .. math:: \partial_t u + a_x \partial_x u - \sigma_x \partial_{xx} u - \sigma_v \partial_{vv} u = 0 \quad \text{in } \Omega_x \times \Omega_v \times (0, T) \\ u = u_0 \quad \text{on } \Omega_x \times \Omega_v \times \{0\} where :math:`u: \Omega_x \times \Omega_v \times (0, T) \to \mathbb{R}` is the unknown function on phase space :math:`(x, v)`, :math:`a_x` is a constant advection velocity in :math:`x` (there is no advection in :math:`v`), and :math:`\sigma_x, \sigma_v` are constant, independent, isotropic diffusion coefficients in :math:`x` and :math:`v` respectively. The initial condition is a separable Gaussian bump. Because advection and diffusion are both constant-coefficient here, the exact solution stays an (un-normalized) Gaussian for all time, with each direction's variance growing linearly under its own diffusion coefficient and its own amplitude decaying to conserve mass -- the classical heat-kernel scaling, applied independently in :math:`x` and :math:`v`: .. math:: u(t, x, v) = \frac{\sigma_{x,0}}{\sigma_x(t)} \frac{\sigma_{v,0}}{\sigma_v(t)} \exp\!\left(-\frac{(x - x_0 - a_x t)^2}{2 \sigma_x(t)^2}\right) \exp\!\left(-\frac{(v - v_0)^2}{2 \sigma_v(t)^2}\right), \quad \sigma_x(t)^2 = \sigma_{x,0}^2 + 2 \sigma_x t, \quad \sigma_v(t)^2 = \sigma_{v,0}^2 + 2 \sigma_v t This example demonstrates `NeuralSemiLagrangian`'s independent `diffusion_coefficient_x`/`diffusion_coefficient_v` arguments (see `advection_diffusion_nd.py` for x-only diffusion) by solving the same problem four times, toggling each coefficient on and off, and comparing the four solutions against their (distinct) exact solutions: - no diffusion: the bump translates in x, unchanged in shape. - diffusion in x only: the bump also spreads and flattens along x. - diffusion in v only: the bump also spreads and flattens along v. - diffusion in both: the bump spreads and flattens along both directions. The domain is periodic (as in `advection_diffusion_nd.py`), and sized with a large margin around the bump's final extent -- even though the exact solution above is the free-space (non-periodic) one, not the RK4/exact-foot path but the *training* itself needs this: collocation points are resampled uniformly over the whole domain at every step, so a domain barely larger than the bump would let a persistent fraction of points near the boundary get shifted outside the network's trained support by the semi-Lagrangian foot at every step. With `periodic=False` (or too tight a margin), that extrapolation noise gets fit into the target at every step and compounds over many steps -- the error stops shrinking with more time steps, or even grows, well before reaching the discretization's true O(dt) regime. With enough margin (here, domain half-width >> bump spread + `VELOCITY_X * END_TIME`), periodic wrap never triggers for any point with non-negligible mass, so it costs nothing relative to the free-space solution while eliminating the extrapolation entirely. """ 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 ( # noqa: E501 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, ) N_COLLOC = 1500 N_EPOCHS_INIT = 200 N_EPOCHS = 10 HIDDEN = [16] * 4 NT = 32 SEED = 0 DIFFUSION = 0.1 # sigma_x / sigma_v when diffusion is "on" for that component VELOCITY_X = 1.0 # constant advection speed in x; none in v X0, V0 = -1.0, 0.0 # initial bump center SIGMA_X0, SIGMA_V0 = 0.4, 0.4 # initial bump width (std) in x and v # Training/sampling domain: wide enough that periodic wrap never triggers for # any point with non-negligible mass (see module docstring). The bump's final # spread is at most sqrt(SIGMA_X0**2 + 2 * DIFFUSION * END_TIME) ~= 0.75, and # its center travels VELOCITY_X * END_TIME = 2 in x; +-6 leaves a wide margin. X_BOUNDS = (-6.0, 6.0) V_BOUNDS = (-6.0, 6.0) DOMAIN_BOUNDS = ((X_BOUNDS[0],), (X_BOUNDS[1],)) DOM_X = Segment1D(X_BOUNDS, is_main_domain=True) DOM_V = Segment1D(V_BOUNDS) SAMPLER = TensorizedSampler( [DomainSampler(DOM_X), UniformVelocitySamplerOnCuboid(DOM_V)], model_type="x_v" ) # Tighter window used only for plotting, since the bump never leaves it. PLOT_X_BOUNDS = (-3.0, 3.0) PLOT_V_BOUNDS = (-3.0, 3.0) END_TIME = 2.0 DOM_T = (0.0, END_TIME) DT = (DOM_T[1] - DOM_T[0]) / NT # (label, diffusion_coefficient_x, diffusion_coefficient_v) CASES = [ ("no diffusion", 0.0, 0.0), ("diffusion in x only", DIFFUSION, 0.0), ("diffusion in v only", 0.0, DIFFUSION), ("diffusion in x and v", DIFFUSION, DIFFUSION), ] def f_init(x: jnp.ndarray, v: jnp.ndarray) -> jnp.ndarray: """Initial condition: a separable Gaussian bump centered at (X0, V0).""" x = jnp.atleast_2d(x) v = jnp.atleast_2d(v) gx = jnp.exp(-((x[:, 0] - X0) ** 2) / (2 * SIGMA_X0**2)) gv = jnp.exp(-((v[:, 0] - V0) ** 2) / (2 * SIGMA_V0**2)) result = gx * gv return result if result.shape[0] == 1 else result.reshape(-1, 1) def make_exact_sol(diffusion_x: float, diffusion_v: float): """Build the exact solution for a given pair of diffusion coefficients. See the module docstring for the closed-form (heat-kernel scaling) formula this implements. """ def exact_sol(t: jnp.ndarray, x: jnp.ndarray, v: jnp.ndarray) -> jnp.ndarray: x = jnp.atleast_2d(x) v = jnp.atleast_2d(v) t = jnp.atleast_1d(t).reshape(x.shape[0], 1)[:, 0] sigma_x_t = jnp.sqrt(SIGMA_X0**2 + 2 * diffusion_x * t) sigma_v_t = jnp.sqrt(SIGMA_V0**2 + 2 * diffusion_v * t) amplitude = (SIGMA_X0 / sigma_x_t) * (SIGMA_V0 / sigma_v_t) gx = jnp.exp(-((x[:, 0] - X0 - VELOCITY_X * t) ** 2) / (2 * sigma_x_t**2)) gv = jnp.exp(-((v[:, 0] - V0) ** 2) / (2 * sigma_v_t**2)) result = amplitude * gx * gv return result if result.shape[0] == 1 else result.reshape(-1, 1) return exact_sol def advection_field( t: jnp.ndarray, x: jnp.ndarray, v: jnp.ndarray ) -> tuple[jnp.ndarray, jnp.ndarray]: """Advection field (a_x, a_v) = (VELOCITY_X, 0): pure translation in x.""" return jnp.ones_like(x) * VELOCITY_X, jnp.zeros_like(v) def exact_characteristic_foot( t: jnp.ndarray, x: jnp.ndarray, v: jnp.ndarray, dt: float ) -> tuple[jnp.ndarray, jnp.ndarray]: """Exact characteristic foot for constant advection: (x - a_x dt, v).""" return x - VELOCITY_X * dt, v def run_case(diffusion_x: float, diffusion_v: float) -> NeuralSemiLagrangian: """Initialize and solve the NSL scheme for one (diffusion_x, diffusion_v) case.""" key = jax.random.PRNGKey(SEED) nn = MLP(in_size=2, out_size=1, hidden_sizes=HIDDEN, key=key) space = ApproximationSpace( {"x": 1, "v": 1}, [(nn, "scalar", None)], model_type="x_v" ) 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=1, model_type="x_v", exact_solution=make_exact_sol(diffusion_x, diffusion_v), diffusion_coefficient_x=diffusion_x, diffusion_coefficient_v=diffusion_v, ) key, neural_sl = neural_sl.initialize( key, space, f_init, N_EPOCHS_INIT, N_COLLOC, file_name="nsl_ad_1d_1v_init", retrain=False, ) key, neural_sl = neural_sl.solve(key, neural_sl.space, N_EPOCHS, N_COLLOC) return neural_sl if __name__ == "__main__": results = [] for label, diffusion_x, diffusion_v in CASES: print(f"Solving case: {label} (sigma_x={diffusion_x}, sigma_v={diffusion_v})") start = time.perf_counter() neural_sl = run_case(diffusion_x, diffusion_v) elapsed = time.perf_counter() - start errors = neural_sl.errors_over_time print( f" ... done in {elapsed:.2f}s, final relative L2 error: " f"{errors[-1, 0]:.2e}" ) results.append((label, diffusion_x, diffusion_v, neural_sl)) # Evaluate every case's trained solution on a common (x, v) grid at t=T # (zoomed into PLOT_X_BOUNDS/PLOT_V_BOUNDS, not the full training domain). n_grid = 200 x_grid = jnp.linspace(*PLOT_X_BOUNDS, n_grid) v_grid = jnp.linspace(*PLOT_V_BOUNDS, n_grid) x_mesh, v_mesh = jnp.meshgrid(x_grid, v_grid) xv_plot = jnp.stack([x_mesh.ravel(), v_mesh.ravel()], axis=1) x_plot, v_plot = xv_plot[:, :1], xv_plot[:, 1:] fig, ax = plt.subplots(2, 2, figsize=(10, 10)) axes = ax.ravel() for (label, diffusion_x, diffusion_v, neural_sl), a in zip(results, axes): space = neural_sl.space 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, v_plot)) u_pred = u_pred.reshape(n_grid, n_grid) cf = a.contourf(x_mesh, v_mesh, u_pred, cmap="turbo", levels=100) plt.colorbar(cf, ax=a) final_l2 = neural_sl.errors_over_time[-1, 0] a.set_title(f"{label}\n(final rel. L2 error: {final_l2:.2e})") a.set_xlabel("x") a.set_ylabel("v") fig.suptitle(f"1D-1V advection-diffusion at t={END_TIME}: effect of diffusion") fig.tight_layout() plt.show()