r"""Solves the Vlasov-Poisson bump-on-tail instability in 1D-1V using the Neural Semi-Lagrangian (NSL) predictor-corrector scheme. .. math:: \partial_t f(t, x, v) + v \partial_x f(t, x, v) + E(t, x) \partial_v f(t, x, v) = 0, \quad \forall (t, x, v) \in [0, T] \times \Omega_x \times \Omega_v, \Delta \Psi(t, x) = \langle f \rangle(t) - \int_{\Omega_v} f(t, x, v) \, dv, \quad E(t, x) = -\nabla_x \Psi(t, x), with :math:`\langle f \rangle(t) = \frac{1}{|\Omega_x|} \int_{\Omega_x} \int_{\Omega_v} f(t, x, v) \, dv \, dx`. Note: this is the sign that makes the coupling *repulsive* (matching the "+E" in the Vlasov advection term above, and the classical Gauss's-law form :math:`-\Delta \Psi = \rho - \langle f \rangle`, as also used in ``src/applications/experimental/scimba_plasma/vlasov_poisson/classical_sl``). The opposite sign was checked numerically (see ``run_reference_sl`` below) and produces unbounded/runaway growth instead of the expected saturating bump-on-tail instability. At every macro time step :math:`[t^n, t^{n+1}]` the coupled system is solved with the predictor-corrector algorithm below, using :class:`NeuralSemiLagrangian` for both transport substeps and a standard elliptic PINN (:class:`Projector` + :class:`LaplacianDirichletND`) for the two Poisson solves: 1. solve :math:`\Delta \Psi^n = \langle f^n \rangle - \int_{\Omega_v} f^n \, dv` to get :math:`E^n = -\nabla_x \Psi^n`; 2. solve the transport equation with :math:`E^n`, from :math:`f^n`, to get a prediction :math:`f^{*}`; 3. solve :math:`\Delta \Psi^{*} = \langle f^{*} \rangle - \int_{\Omega_v} f^{*} \, dv` to get :math:`E^{*}`, and set :math:`\bar{E} = (E^n + E^{*}) / 2`; 4. solve the transport equation with :math:`\bar{E}`, from :math:`f^n` (not :math:`f^{*}`), to get :math:`f^{n+1}`. The initial condition is the classical bump-on-tail distribution: a Maxwellian core plus a warmer, drifting "bump" in velocity, with a small cosine density perturbation in :math:`x` that seeds the instability. """ # %% import os import time import warnings 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.linear_approximation.basis.analytic_bases import local_lagrange_basis from scimba_jax.linear_approximation.basis.general_bases import AnalyticBasis from scimba_jax.linear_approximation.galerkin.fem.elliptic_fe_scheme import ( EllipticFEscheme, ) from scimba_jax.linear_approximation.meshes.mesh import Mesh from scimba_jax.linear_approximation.quad.gauss_quad import UnitSquareTensorized from scimba_jax.linear_approximation.variables.variables_fe import VariablesFE from scimba_jax.mapping.mapping import InvertibleFunction, Mapping 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.networks.mlp import MLP from scimba_jax.nonlinear_approximation.numerical_solvers.neural_semi_lagrangian import ( NeuralSemiLagrangian, ) from scimba_jax.nonlinear_approximation.numerical_solvers.projectors import Projector from scimba_jax.physical_models.abstract_physical_weak_model import ( AbstractPhysicalWeakModel, ) from scimba_jax.physical_models.classical_weakform.laplacian_weak_form import ( LaplacianWeakForm, ) from scimba_jax.physical_models.elliptic_pde.laplacians import LaplacianDirichletND FIG_DIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), "fig") os.makedirs(FIG_DIR, exist_ok=True) # ───────────────────────────────────────────────────────────────────────────── # Physical parameters of the bump-on-tail test case # ───────────────────────────────────────────────────────────────────────────── X_MIN, X_MAX = 0.0, 10 * jnp.pi V_MIN, V_MAX = -9.0, 9.0 SHIFT = 0.0 # shift the distribution function to avoid negative values EPS_BOT = 0.1 # fraction of the population in the "bump" (tail) T_EQ = 1.0 # temperature of the Maxwellian bulk T_BOT = 0.2 # temperature of the bump V_BOT = 3.8 # drift velocity of the bump EPSILON = 0.03 # amplitude of the initial density perturbation KX = 0.4 # wavenumber of the initial density perturbation DOMAIN_X = Segment1D((X_MIN, X_MAX), is_main_domain=True) DOMAIN_V = Segment1D((V_MIN, V_MAX)) DOMAIN_BOUNDS = ( (X_MIN, V_MIN), # lower bounds for (x, v) (X_MAX, V_MAX), # upper bounds for (x, v) ) # domain/sampler for the electric potential (1D in x only) DOMAIN_X_PHI = Segment1D((X_MIN, X_MAX), is_main_domain=True) SAMPLER_PHI = TensorizedSampler([DomainSampler(DOMAIN_X_PHI)], model_type="x") # Sampler for phase space (x, v) SAMPLER = TensorizedSampler( [ DomainSampler(DOMAIN_X), UniformVelocitySamplerOnCuboid(DOMAIN_V), ], model_type="x_v", ) # Time domain and macro time-stepping DOM_T = (0.0, 12.0) NT = int(DOM_T[1] - DOM_T[0]) * 5 # number of macro (predictor-corrector) time steps DT = (DOM_T[1] - DOM_T[0]) / NT # macro time step size N_RK_SUBSTEPS = 1 # RK4 substeps per NSL transport solve # Network sizes HIDDEN_F = [30] * 6 HIDDEN_PHI = [20] * 6 # Training hyperparameters N_COLLOC_F = 64**2 # collocation points for the (x, v) phase space N_COLLOC_PHI = 256 # collocation points for the electric potential N_QUAD_V = 128 # velocity quadrature points used to compute the charge density N_EPOCHS_INIT = 250 # epochs to fit the initial condition N_EPOCHS_F = 100 # epochs per Vlasov transport substep N_EPOCHS_PHI_INIT = 350 # epochs for the very first Poisson solve N_EPOCHS_PHI = 60 # epochs for subsequent (warm-started) Poisson solves # 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, } # ───────────────────────────────────────────────────────────────────────────── # Poisson solver: three interchangeable ways to get Phi (hence E) from f, all # behind the same two-function interface -- solve_poisson(key, phi_state, # space_f, n_epochs, ...) -> (key, phi_state) and make_electric_field(phi_state) # -> (pointwise callable x -> E(x)) -- selected below and bound once, so # `time_step` and friends never branch on POISSON_SOLVER themselves: # - "pinn": the original elliptic PINN (Projector + LaplacianDirichletND), # periodic by construction (periodic embedding), warm-started/cached like f. # - "fft": direct spectral solve (rho -> FFT -> divide by ik -> E_hat), # exact on a periodic domain, no training; see solve_poisson_fft below. # - "fem": direct Galerkin solve (LaplacianWeakForm + EllipticFEscheme), # reusing the classical-approach FEM infra as-is; see solve_poisson_fem. POISSON_SOLVER = "pinn" # "pinn" | "fft" | "fem" # grid resolution for the FFT-based Poisson solve (see solve_poisson_fft) N_X_POISSON_FFT = 1024 # mesh/basis for the FEM-based Poisson solve (see solve_poisson_fem) N_CELLS_PHI_FEM = 256 FEM_BASIS_ORDER = 3 FEM_QUAD_ORDER = 6 # Periodic embedding in x: the domain length X_MAX - X_MIN = 10*pi is exactly two # wavelengths of the seed perturbation (2*pi/KX = 5*pi), i.e. its fundamental is the # *second* harmonic of the domain-periodic basis cos/sin(n * 2*pi*x/(X_MAX-X_MIN)). # A handful more harmonics are kept to resolve the finer structure (filamentation) # the instability develops over time. Phi is the (smoothed, Poisson-filtered) integral # of f, so it needs far fewer harmonics than f itself to be represented accurately. X_PERIOD = X_MAX - X_MIN N_PERIODIC_FEATURES_F = 16 N_PERIODIC_FEATURES_PHI = 8 # Reference (classical, grid-based) semi-Lagrangian solver, used as ground truth to # check the NSL solution against. Strang splitting + FFT: the x- and v-advection # substeps are exact spectral shifts, and the Poisson equation is solved spectrally. # dt_ref is chosen independently of the NSL macro step DT -- this is a different # numerical scheme with its own splitting-error budget. N_X_REF = 1024 N_V_REF = 512 DT_REF = 0.01 # The v-advection substep is an FFT phase shift, which is exact only on a *periodic* # grid -- but v is not physically periodic, so any density that would drift past # V_MIN/V_MAX must not be allowed to wrap back in from the other side. We therefore # advect v on a padded grid (V_PAD wider on each side) and zero out that padding after # every substep, so mass that exits (V_MIN, V_MAX) is dropped rather than aliased back # in -- the free-streaming ("open boundary") analogue of a backward semi-Lagrangian # foot landing outside the resolved domain, instead of a periodic one wrapping around. V_PAD = 3.0 def f_init(x: jnp.ndarray, v: jnp.ndarray) -> jnp.ndarray: """Bump-on-tail initial condition. Args: x: Spatial points of shape (1,). v: Velocity points of shape (1,). Returns: Initial values of shape (1,). """ factor1 = (1 - EPS_BOT) / jnp.sqrt(2 * jnp.pi * T_EQ) factor2 = 2 * EPS_BOT / jnp.sqrt(2 * jnp.pi * T_BOT) f1_eq = factor1 * jnp.exp(-0.5 * v**2 / T_EQ) f2_eq = factor2 * jnp.exp(-((v - V_BOT) ** 2) / (2 * T_BOT)) feq = f1_eq + f2_eq return SHIFT + feq * (1 + EPSILON * jnp.cos(KX * x)) def _f_on_grid(space_f, x_grid: jnp.ndarray, v_grid: jnp.ndarray) -> jnp.ndarray: """Evaluate f on the tensor grid (x_grid, v_grid); shape (len(x_grid), len(v_grid)). Shared by every consumer that needs f on a regular (x, v) grid: the PINN Poisson RHS quadrature (compute_f_average, make_rho_rhs) and the FFT Poisson solve (solve_poisson_fft) all reduce to this same evaluation. """ variables_f = space_f.create_variables() f_batched = variables_f[0].vmap_on_physical_variables() xx, vv = jnp.meshgrid(x_grid, v_grid, indexing="ij") f_vals = f_batched(space_f, xx.reshape(-1, 1), vv.reshape(-1, 1)) return f_vals.reshape(x_grid.shape[0], v_grid.shape[0]) def compute_f_average(space_f, n_x: int = 256, n_v: int = N_QUAD_V) -> jnp.ndarray: """Compute = (1 / |Omega_x|) * int_x int_v f dv dx on a regular grid.""" x_grid = jnp.linspace(X_MIN, X_MAX, n_x, endpoint=False) v_grid = jnp.linspace(V_MIN, V_MAX, n_v, endpoint=False) f_vals = _f_on_grid(space_f, x_grid, v_grid) rho = jnp.sum(f_vals, axis=1) * (V_MAX - V_MIN) / n_v return jnp.mean(rho) def make_rho_rhs(space_f, f_average: jnp.ndarray, n_v_quad: int = N_QUAD_V): """Build the (pointwise) RHS x -> rho(x) - of the Poisson residual. rho(x) = int_v f(x, v) dv is computed by rectangle quadrature in v, evaluating the current f approximation space. The sign matches LaplacianResidual's convention (LHS = -Delta Phi), so LHS = RHS enforces -Delta Phi = rho - , i.e. the standard (repulsive) Gauss's-law form: Delta Phi = - rho, E = -grad Phi. See the module docstring / run_reference_sl for why (contrary to a first, literal reading of Delta Psi = rho - with E = -grad Psi) this is the sign that actually reproduces the bump-on-tail instability instead of an unbounded, attractive-self-coupling runaway. """ variables_f = space_f.create_variables() f_batched = variables_f[0].vmap_on_physical_variables() v_quad = jnp.linspace(V_MIN, V_MAX, n_v_quad, endpoint=False)[:, None] dv = (V_MAX - V_MIN) / n_v_quad def rho_rhs(x: jnp.ndarray) -> jnp.ndarray: x_rep = jnp.broadcast_to(x, (n_v_quad, x.shape[-1])) f_vals = f_batched(space_f, x_rep, v_quad)[:, 0] rho = jnp.sum(f_vals) * dv return jnp.array([rho - f_average]) return rho_rhs def solve_poisson_pinn( model_phi, key, space_phi, space_f, n_epochs: int, n_colloc: int, file_name: str | None = None, retrain: bool = True, ): """One elliptic PINN solve of Delta Phi = - rho[f], warm-started from space_phi. ``model_phi`` must be a :class:`LaplacianDirichletND` built once, outside of any ``jax.lax.scan``: constructing a fresh :class:`AbstractPhysicalModel` triggers a concrete-only boundary computation on ``main_domain`` (``main_domain. get_all_boundaries()``, which calls ``.item()`` on array bounds) that is incompatible with jax's tracing. Instead, we reuse one pre-built model and just mutate its ``f_rhs`` -- the same pattern :class:`NeuralSemiLagrangian` itself uses for its own (scan-nested) target function. """ f_average = compute_f_average(space_f) rho_rhs = make_rho_rhs(space_f, f_average) model_phi.physical_residuals[DOMAIN_X_PHI.get_label()].f_rhs = rho_rhs proj_phi = Projector(model_phi, space_phi, SAMPLER_PHI, **ENG_KWARGS) if (file_name is not None) and (not retrain): # Load if we should load n_epochs_load, proj_phi = proj_phi.load(file_name) # If the projector is loaded, return its space if n_epochs_load > 0: print(f"Loaded pre-trained model from {file_name}.") return proj_phi.key, proj_phi.space # Otherwise, train, and save if needed key, proj_phi = proj_phi.project_scan(key, space_phi, n_epochs, n_colloc) if file_name is not None: proj_phi.save(file_name) return key, proj_phi.space def make_electric_field_pinn(space_phi): """Build a pointwise x -> E(x) = -grad_x Phi(x) function from a potential space. Pointwise (single point in, single point out): this is how :class:`Characteristic` calls the advection field (it is vmapped/batched from the outside, as an RHS of the NSL transport residual). Use ``jax.vmap`` at the call site when a batched evaluation is needed (e.g. for diagnostics/plots). """ variables_phi = space_phi.create_variables() minus_grad_phi = -variables_phi[0].gradient("x") def electric_field(x: jnp.ndarray) -> jnp.ndarray: return minus_grad_phi(space_phi, x) return electric_field def make_advection_field(space_phi): """Advection field (v, E(x)) for the Vlasov characteristics, from one potential.""" electric_field = make_electric_field(space_phi) def advection_field(t, x, v): return v, electric_field(x) return advection_field def make_averaged_advection_field(space_phi_a, space_phi_b): """Advection field (v, (E_a(x) + E_b(x)) / 2), for the corrector substep.""" field_a = make_electric_field(space_phi_a) field_b = make_electric_field(space_phi_b) def advection_field(t, x, v): return v, 0.5 * (field_a(x) + field_b(x)) return advection_field # ───────────────────────────────────────────────────────────────────────────── # Alternate Poisson solver #1: direct spectral (FFT) solve. # # No training, no warm start needed -- each call recomputes rho on a grid (same # quadrature-over-v pattern as the PINN RHS) and solves spectrally, reusing # ``_charge_and_field`` (the very formula ``run_reference_sl`` uses every # substep). ``phi_state`` here is just the raw array of Fourier coefficients of # E, so it is evaluated off-grid by summing its (truncated) Fourier series -- # the FFT counterpart of "re-interpolating to create a function" mentioned for # the torch version of this example, except here the interpolation is exact # (band-limited, matching the spectral solve) rather than a spline fit. # ───────────────────────────────────────────────────────────────────────────── def solve_poisson_fft( space_f, n_x: int = N_X_POISSON_FFT, n_v: int = N_QUAD_V ) -> jnp.ndarray: """Direct spectral solve of Delta Phi = - rho[f]; returns E's FFT coefficients.""" x_grid = jnp.linspace(X_MIN, X_MAX, n_x, endpoint=False) v_grid = jnp.linspace(V_MIN, V_MAX, n_v, endpoint=False) dv = (V_MAX - V_MIN) / n_v f_grid = _f_on_grid(space_f, x_grid, v_grid) kx = jnp.fft.fftfreq(n_x) * n_x * 2.0 * jnp.pi / X_PERIOD _, e_grid = _charge_and_field(f_grid, dv, kx) return jnp.fft.fft(e_grid) def _fourier_eval( x: jnp.ndarray, values_hat: jnp.ndarray, kx: jnp.ndarray ) -> jnp.ndarray: """Band-limited (trigonometric) interpolation of a periodic grid function at x. ``values_hat`` are the FFT coefficients of a function sampled on the ``X_PERIOD``-periodic grid starting at ``X_MIN``; this evaluates the exact inverse-DFT formula off-grid, i.e. the unique trigonometric polynomial (through the same modes ``kx``) that interpolates the grid samples. """ n = values_hat.shape[-1] phase = jnp.exp(1j * kx * (x - X_MIN)) return jnp.real(jnp.sum(values_hat * phase)) / n def make_electric_field_fft(e_hat: jnp.ndarray, n_x: int = N_X_POISSON_FFT): """Pointwise x -> E(x) from the FFT coefficients returned by solve_poisson_fft.""" kx = jnp.fft.fftfreq(n_x) * n_x * 2.0 * jnp.pi / X_PERIOD def electric_field(x: jnp.ndarray) -> jnp.ndarray: return jnp.array([_fourier_eval(x[0], e_hat, kx)]) return electric_field # ───────────────────────────────────────────────────────────────────────────── # Alternate Poisson solver #2: classical FEM (Galerkin), reusing the # LaplacianWeakForm / EllipticFEscheme machinery from # examples/examples_jax/fem/solve/classical_approach/solve_1d_laplacian.py # as-is -- no FEM code is reimplemented here, only the problem-specific RHS # (make_rho_rhs, shared with the PINN solver) and mesh/mapping/BC wiring. # # No training needed -- a Galerkin solve of a linear problem converges in a # single Newton step -- but a persistent, mutated scheme object *is* still # needed across scan iterations, for the same reason as the PINN model_phi # (see solve_poisson_pinn): ``EllipticFEscheme.__init__`` computes and caches # the Dirichlet boundary node *positions* via the mesh mapping, which # `elliptic_fe_scheme.py` documents as requiring a concrete (non-traced) mesh # mapping ("a tracer under JIT and could not be turned into numpy there"). # So build_fem_scheme() builds the Mesh/basis/scheme once, outside any # jax.lax.scan, and solve_poisson_fem only ever mutates its RHS field and # re-solves -- never reconstructs the scheme. # # Boundary condition: homogeneous Dirichlet Phi = 0 at both x = X_MIN and # x = X_MAX. There is no periodic BC in the FEM infra, but this is not the # same trap the PINN docstring warns about for a *soft*, loss-penalized # Dirichlet anchor: here the Dirichlet condition is imposed exactly (by # lifting, not by a training penalty), and int(rho - ) dx = 0 *exactly* by # construction of -- so E(X_MIN) = E(X_MAX) automatically (any two # particular solutions of a 1D Poisson problem differ by an affine function, # whose derivative is a constant that cancels out of E(X_MAX) - E(X_MIN)). # The Dirichlet-0 solve therefore *is* the periodic solution, up to the # Phi = 0 gauge (irrelevant since only E = -grad Phi is ever used). # ───────────────────────────────────────────────────────────────────────────── def _fem_mapping() -> Mapping: """Affine map from the reference cell [0, 1] to the physical [X_MIN, X_MAX].""" def fwd(xi: jnp.ndarray) -> jnp.ndarray: return jnp.array([X_MIN + X_PERIOD * xi[0]]) def inv(x: jnp.ndarray) -> jnp.ndarray: return jnp.array([(x[0] - X_MIN) / X_PERIOD]) return Mapping(mappings=[InvertibleFunction(fwd, inv)]) def build_fem_scheme( n_cells: int = N_CELLS_PHI_FEM, order: int = FEM_BASIS_ORDER, quad_order: int = FEM_QUAD_ORDER, ) -> EllipticFEscheme: """Build the (mutable, reused) FEM scheme once, outside of any jax.lax.scan. See the module comment above solve_poisson_fem for why this must not be rebuilt inside a scanned step. ``f=None`` is a placeholder: solve_poisson_fem overwrites it with the current RHS before every solve, exactly how solve_poisson_pinn mutates ``model_phi``'s ``f_rhs``. """ mesh = Mesh( dim=1, n_cells=[n_cells], ref_quad=UnitSquareTensorized(dim=1, order=quad_order), mapping=_fem_mapping(), ) basis = AnalyticBasis( nb_basis=order + 1, out_dim=1, mesh=mesh, local_basis=lambda coords, i, m, _order=order: local_lagrange_basis( coords, i, m, order=_order, out_dim=1 ), basis_type="scalar", ) variables = VariablesFE(basis=basis, nb_variables=1) weak_form = LaplacianWeakForm(dim=1, f=None) return EllipticFEscheme( AbstractPhysicalWeakModel.from_weak_form( weak_form, dirichlet=lambda x: jnp.zeros(1) ), variables, ) def solve_poisson_fem( fem_scheme: EllipticFEscheme, space_f, n_v_quad: int = N_QUAD_V ) -> VariablesFE: """One direct FEM (Galerkin) solve of Delta Phi = - rho[f]. ``fem_scheme`` must be built once by :func:`build_fem_scheme`, outside of any ``jax.lax.scan`` -- this only mutates its RHS field and re-solves. """ f_average = compute_f_average(space_f, n_v=n_v_quad) rho_rhs = make_rho_rhs(space_f, f_average, n_v_quad) fem_scheme.pde.weak_forms["main"].f = rho_rhs solved_scheme = EllipticFEscheme.solve(fem_scheme) return solved_scheme.variables def make_electric_field_fem(variables: VariablesFE): """Pointwise x -> E(x) = -grad_x Phi(x) from a solved FEM potential. ``VariablesFE._classical_local_evaluate_pure`` is the *differentiable* pointwise-evaluate (unlike ``variables.evaluate``/``.local_evaluate``, documented as non-differentiable "uses self" shortcuts), so autodiff through it gives the FE gradient directly -- no need to hand-assemble it from ``trial_basis.derivative`` and the local DOFs. """ def phi(x: jnp.ndarray) -> jnp.ndarray: return VariablesFE._classical_local_evaluate_pure(variables, variables.dofsl, x) grad_phi = jax.jacobian(phi) def electric_field(x: jnp.ndarray) -> jnp.ndarray: return -grad_phi(x)[:, 0] return electric_field def _charge_and_field(f: jnp.ndarray, dv: float, kx: jnp.ndarray) -> tuple: """rho(x) = int_v f dv, and E = -grad_x Phi with Delta Phi = - rho (spectral). With f(x) = sum_k f_hat(k) exp(ikx): -Delta Phi = rho - gives k^2 Phi_hat(k) = rho_hat(k) for k != 0 (the k=0 mode of the RHS is identically zero since is subtracted, so Phi_hat(0) is gauge and set to 0), and E = -Phi' gives E_hat(k) = -(ik) Phi_hat(k) = -i * rho_hat(k) / k. This is the standard (repulsive) Gauss's-law sign -- see the module docstring -- and matches ``VlasovPoisson.compute_efield`` in ``src/applications/experimental/scimba_plasma/vlasov_poisson/classical_sl/vlasov.py``. """ rho = jnp.sum(f, axis=1) * dv rho_hat = jnp.fft.fft(rho) kx_safe = jnp.where(kx != 0, kx, 1.0) E_hat = jnp.where(kx != 0, -1j * rho_hat / kx_safe, 0.0 + 0.0j) E = jnp.fft.ifft(E_hat).real return rho, E def _bspline(p: int, j: int, x: jnp.ndarray) -> jnp.ndarray: """Value at x in [0, 1] of the degree-p B-spline with support starting at node j. De Boor's recursion; p and j are static Python ints (only x may be a JAX array), so this unrolls into a fixed-depth (= p) computation graph under jit/vmap. Ported from ``bspline`` in ``src/applications/experimental/scimba_plasma/vlasov_poisson/classical_sl/bsl.py``. """ if p == 0: return jnp.ones_like(x) if j == 0 else jnp.zeros_like(x) w = (x - j) / p w1 = (x - j - 1) / p return w * _bspline(p - 1, j, x) + (1 - w1) * _bspline(p - 1, j + 1, x) def _bspline_eigenvalues(p: int, ncells: int) -> tuple[jnp.ndarray, jnp.ndarray]: """Precompute the (static) FFT modes and B-spline collocation-matrix eigenvalues. Ported from ``BSpline.__init__`` in the classical_sl reference; depends only on (p, ncells), so it is computed once per grid and reused every step. """ modes = 2.0 * jnp.pi * jnp.arange(ncells) / ncells eig_bspl = _bspline(p, -(p + 1) // 2, 0.0) * jnp.ones_like(modes) for j in range(1, (p + 1) // 2): eig_bspl = eig_bspl + _bspline(p, j - (p + 1) // 2, 0.0) * 2.0 * jnp.cos( j * modes ) return modes, eig_bspl def _bspline_shift( f: jnp.ndarray, alpha: jnp.ndarray, p: int, deltax: float, modes: jnp.ndarray, eig_bspl: jnp.ndarray, ) -> jnp.ndarray: """Periodic degree-p spline-interpolated shift f(x) -> f(x - alpha), on a 1D array. Exact spline interpolation via the circulant-matrix/FFT trick (no local stencil gather needed): ported from ``BSpline.interpolate_disp`` in the classical_sl reference. f is a single 1D row/column (shape (ncells,)); vmap this over the other axis, pairing each row/column with its own (per-index) shift amount alpha. """ ncells = f.shape[-1] ishift = jnp.floor(-alpha / deltax) beta = -ishift - alpha / deltax eigalpha = jnp.zeros(ncells, dtype=complex) for j in range(-(p - 1) // 2, (p + 1) // 2 + 1): eigalpha = eigalpha + _bspline(p, j - (p + 1) // 2, beta) * jnp.exp( (ishift + j) * 1j * modes ) return jnp.real(jnp.fft.ifft(jnp.fft.fft(f) * eigalpha / eig_bspl)) N_BSPLINE_DEGREE = 3 # cubic, as in the classical_sl reference def run_reference_sl( t_final: float, nx: int = N_X_REF, nv: int = N_V_REF, dt_ref: float = DT_REF, v_pad: float = V_PAD, ) -> tuple[jnp.ndarray, jnp.ndarray, jnp.ndarray, jnp.ndarray, jnp.ndarray]: """Classical (grid-based) semi-Lagrangian Vlasov-Poisson solver, jitted. Cheng-Knorr / Sonnendrucker-style Strang splitting with cubic B-spline advection, matching ``VlasovPoisson`` in ``src/applications/experimental/scimba_plasma/vlasov_poisson/classical_sl``: advect x by dt/2 (periodic cubic-spline shift, dx/dt = v) -> solve Poisson for E(x) (spectral) -> advect v by dt (cubic-spline shift, dv/dt = E) -> advect x by dt/2. x is genuinely periodic, so its spline shift is exact as-is (matching the reference exactly). v is not physically periodic, so unlike the reference it is advected on a grid padded by v_pad on each side; that padding is zeroed out after every v-substep, so any density that exits (V_MIN, V_MAX) is dropped instead of wrapping back in from the other side. f_final is truncated back to the physical (V_MIN, V_MAX) grid before being returned. Args: t_final: final time to integrate to. nx: number of grid points in x. nv: number of grid points in v, on the physical (V_MIN, V_MAX) grid. dt_ref: (fixed) time step of the reference solver, independent of the NSL dt. v_pad: width of the (discarded) velocity buffer on each side of (V_MIN, V_MAX). Returns: (f_final, x_grid, v_grid, times, electric_energy_history), where f_final has shape (nx, nv) and electric_energy_history[k] is the energy at times[k]. """ p = N_BSPLINE_DEGREE x_grid = jnp.linspace(X_MIN, X_MAX, nx, endpoint=False) v_grid = jnp.linspace(V_MIN, V_MAX, nv, endpoint=False) dx = X_PERIOD / nx dv = (V_MAX - V_MIN) / nv # padded velocity grid: n_pad extra cells (spacing dv) on each side of (V_MIN, V_MAX) n_pad = int(round(v_pad / dv)) nv_ext = nv + 2 * n_pad v_grid_ext = V_MIN - n_pad * dv + dv * jnp.arange(nv_ext) physical_mask = jnp.zeros(nv_ext, dtype=bool).at[n_pad : n_pad + nv].set(True) xx, vv_ext = jnp.meshgrid(x_grid, v_grid_ext, indexing="ij") f0 = f_init(xx, vv_ext) * physical_mask[None, :] kx = jnp.fft.fftfreq(nx) * nx * 2.0 * jnp.pi / X_PERIOD modes_x, eig_bspl_x = _bspline_eigenvalues(p, nx) modes_v, eig_bspl_v = _bspline_eigenvalues(p, nv_ext) def advect_x(f, dt_sub): # shift each v-column j by v_grid_ext[j] * dt_sub (dx/dt = v) def shift_col(f_col, alpha): return _bspline_shift(f_col, alpha, p, dx, modes_x, eig_bspl_x) return jax.vmap(shift_col, in_axes=(1, 0), out_axes=1)(f, dt_sub * v_grid_ext) def advect_v(f, e_field, dt_sub): # shift each x-row i by e_field[i] * dt_sub (dv/dt = E) def shift_row(f_row, alpha): return _bspline_shift(f_row, alpha, p, dv, modes_v, eig_bspl_v) f = jax.vmap(shift_row, in_axes=(0, 0), out_axes=0)(f, dt_sub * e_field) # drop whatever drifted outside (V_MIN, V_MAX): open, not periodic, boundary return f * physical_mask[None, :] def body(f, _): f = advect_x(f, dt_ref / 2.0) _, e_field = _charge_and_field(f, dv, kx) f = advect_v(f, e_field, dt_ref) f = advect_x(f, dt_ref / 2.0) _, e_field = _charge_and_field(f, dv, kx) energy = 0.5 * jnp.mean(e_field**2) * X_PERIOD return f, energy n_steps = int(round(t_final / dt_ref)) run_scan = jax.jit( lambda f_init_: jax.lax.scan(body, f_init_, xs=None, length=n_steps) ) f_final_ext, energy_history = run_scan(f0) times = (jnp.arange(n_steps) + 1) * dt_ref f_final = f_final_ext[:, n_pad : n_pad + nv] return f_final, x_grid, v_grid, times, energy_history def plot_solution_vs_reference( space_f, f_ref: jnp.ndarray, x_grid: jnp.ndarray, v_grid: jnp.ndarray, t: float ) -> None: """Compare the NSL density against the reference SL solution, on the reference grid. Mirrors the NSL/SL/|NSL - SL| plotting style of linear_vlasov_1d_1v.py. The NSL space is evaluated directly on the reference solver's own (x_grid, v_grid), so no interpolation is needed for the comparison. """ n_x, n_v = x_grid.shape[0], v_grid.shape[0] xx, vv = jnp.meshgrid(x_grid, v_grid, indexing="ij") variables_f = space_f.create_variables() f_batched = variables_f[0].vmap_on_physical_variables() f_nsl = jax.device_get( f_batched(space_f, xx.reshape(-1, 1), vv.reshape(-1, 1)) ).reshape(n_x, n_v) f_ref = jax.device_get(f_ref) error = jnp.abs(f_nsl - f_ref) relative_l2 = jnp.sqrt(jnp.mean(error**2)) / jnp.sqrt(jnp.mean(f_ref**2)) def plot_u(ax, u, kind, has_title=True): cmap = "inferno" if kind == "|NSL - SL|" else "turbo" im = ax.contourf(x_grid, v_grid, u.T, levels=256, cmap=cmap, zorder=-20) plt.colorbar(im, ax=ax) ax.contour( im, levels=im.levels[::32], 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:.2f} [{POISSON_SOLVER}] " f"(min={float(u.min()):.2f}, max={float(u.max()):.2f})" ) else: title = ( rf"{kind}: t={t:.2f} [{POISSON_SOLVER}] " rf"(rel. $L^2$={float(relative_l2):.2e})" ) if has_title: ax.set_title(title) ax.set_xlabel("x") ax.set_ylabel("v") for qty, kind in zip([f_nsl, f_ref, error], ["NSL", "SL", "|NSL - SL|"]): fig, ax = plt.subplots(1, 1, figsize=(5, 4)) plot_u(ax, qty, kind, has_title=False) plt.tight_layout() fig_name = kind.replace(" ", "_").replace("|", "") plt.savefig( os.path.join( FIG_DIR, f"bump_on_tail_{fig_name}_{POISSON_SOLVER}_t{t:.2f}.pdf" ) ) plt.show() print(f"Relative L2 error (NSL vs reference SL) at t={t:.2f}: {relative_l2:.3e}") def plot_loss_history(losses_history: jnp.ndarray) -> None: """Plot the training loss history: initial condition fit + every transport substep. Args: losses_history: 1D array of "total" loss values, in chronological order. """ fig, ax = plt.subplots(1, 1, figsize=(6, 4)) ax.semilogy(losses_history, color="tomato", lw=0.8) ax.set_xlabel("training epoch (initial fit, then every transport substep)") ax.set_ylabel("total loss") ax.set_title(f"bump-on-tail: NSL training loss history [{POISSON_SOLVER}]") plt.tight_layout() plt.savefig( os.path.join(FIG_DIR, f"bump_on_tail_loss_history_{POISSON_SOLVER}.pdf") ) plt.show() def plot_electric_energy( times, energies, ref_times: jnp.ndarray | None = None, ref_energies: jnp.ndarray | None = None, ) -> None: """Plot the electric energy 0.5 * int E^2 dx as a function of time. If a reference (times, energies) curve is given, it is overlaid for comparison -- the classical bump-on-tail growth-rate diagnostic. """ fig, ax = plt.subplots(1, 1, figsize=(6, 4)) if ref_times is not None: ax.semilogy( ref_times, ref_energies, color="steelblue", lw=1.2, label="reference SL" ) ax.semilogy(times, energies, "o", color="tomato", label="NSL") if ref_times is not None: ax.legend() ax.set_xlabel("t") ax.set_ylabel(r"$\frac{1}{2}\int E^2 dx$") ax.set_title(f"bump-on-tail: electric energy growth [{POISSON_SOLVER}]") plt.tight_layout() plt.savefig( os.path.join(FIG_DIR, f"bump_on_tail_electric_energy_{POISSON_SOLVER}.pdf") ) plt.show() # %% if __name__ == "__main__": """Solve the Vlasov-Poisson bump-on-tail instability with Neural Semi-Lagrangian.""" key = jax.random.PRNGKey(0) # distribution function f(x, v): periodic embedding on x only (axis 0), v (axis 1) # stays a plain linear input. key, subkey = jax.random.split(key) nn_f = MLP( in_size=2, out_size=1, activation="sin", activation_output="relu4", hidden_sizes=HIDDEN_F, key=subkey, embedding="periodic", periods=(X_PERIOD,), embedding_axes=(0,), n_periodic_features=N_PERIODIC_FEATURES_F, ) space_f = ApproximationSpace( {"x": 1, "v": 1}, [(nn_f, "scalar", None)], model_type="x_v" ) # `phi_state_n`/`phi_state_star` carry whatever the selected POISSON_SOLVER # returns -- a (periodic-embedding) ApproximationSpace for "pinn", warm-started # across steps; a bare array of E's FFT coefficients for "fft"; a solved # VariablesFE for "fem" -- opaquely to the rest of the script, which only ever # calls them through solve_poisson/make_electric_field (bound below). Only # "pinn" needs a real initial state (as a warm start); the direct solvers # ignore whatever they are handed, so None is enough of a placeholder. if POISSON_SOLVER == "pinn": # Periodic embedding makes Phi exactly periodic by construction, so no # boundary condition/anchor is needed (a Dirichlet Phi=0 anchor at both # ends would only match values there, not derivatives, injecting a # spurious force discontinuity exactly where trajectories wrap around -- # see solve_poisson_fem's docstring for why this trap does *not* apply to # the FEM solver's own, exactly-imposed Dirichlet-0 BC). key, subkey = jax.random.split(key) nn_phi_n = MLP( in_size=1, out_size=1, hidden_sizes=HIDDEN_PHI, key=subkey, embedding="periodic", periods=(X_PERIOD,), embedding_axes=(0,), n_periodic_features=N_PERIODIC_FEATURES_PHI, ) phi_state_n = ApproximationSpace( {"x": 1}, [(nn_phi_n, "scalar", None)], model_type="x" ) key, subkey = jax.random.split(key) nn_phi_star = MLP( in_size=1, out_size=1, hidden_sizes=HIDDEN_PHI, key=subkey, embedding="periodic", periods=(X_PERIOD,), embedding_axes=(0,), n_periodic_features=N_PERIODIC_FEATURES_PHI, ) phi_state_star = ApproximationSpace( {"x": 1}, [(nn_phi_star, "scalar", None)], model_type="x" ) # Build the (mutable, reused) model object once, *outside* the scan below. # LaplacianDirichletND is a plain Python object (not a jax pytree); # constructing it triggers a concrete-only boundary computation on # main_domain that jax's tracing cannot handle, so it must not be rebuilt # inside the scanned time_step. Instead, each step mutates its f_rhs -- # exactly how NeuralSemiLagrangian.solve itself updates its own # (scan-nested) target function. The FFT/FEM solvers need no such # persistent object (see their own docstrings above). model_phi = LaplacianDirichletND(DOMAIN_X_PHI, bc="strong", model_type="x") def solve_poisson( key, phi_state, space_f, n_epochs, file_name=None, retrain=True ): return solve_poisson_pinn( model_phi, key, phi_state, space_f, n_epochs, N_COLLOC_PHI ) make_electric_field = make_electric_field_pinn elif POISSON_SOLVER == "fft": phi_state_n = None phi_state_star = None def solve_poisson( key, phi_state, space_f, n_epochs=None, file_name=None, retrain=True ): return key, solve_poisson_fft(space_f) make_electric_field = make_electric_field_fft elif POISSON_SOLVER == "fem": phi_state_n = None phi_state_star = None # Built once, outside the scan, and only ever mutated + re-solved -- # see the module comment above solve_poisson_fem for why. fem_scheme = build_fem_scheme() def solve_poisson( key, phi_state, space_f, n_epochs=None, file_name=None, retrain=True ): return key, solve_poisson_fem(fem_scheme, space_f) make_electric_field = make_electric_field_fem else: raise ValueError(f"Unknown POISSON_SOLVER: {POISSON_SOLVER!r}") def _placeholder_field(t, x, v): return v, jnp.zeros_like(x) nsl = NeuralSemiLagrangian( main_domain=DOMAIN_X, time_domain=(0.0, DT), sampler=SAMPLER, dt=DT, advection_field=_placeholder_field, periodic=True, domain_bounds=DOMAIN_BOUNDS, out_size=1, model_type="x_v", n_rk_steps=N_RK_SUBSTEPS, ) print("Initializing f^0 (bump-on-tail initial condition)...") start_init = time.perf_counter() key, nsl = nsl.initialize( key, space_f, f_init, N_EPOCHS_INIT, N_COLLOC_F, file_name="bump_on_tail_init", retrain=False, **ENG_KWARGS, ) end_init = time.perf_counter() space_f = nsl.space print(f"Initializing f^0... Done in {end_init - start_init:.2f} seconds\n") # cold-start Poisson solve for E^0 (untrained potential net -> more epochs); # every later E^n solve, inside the scan below, is warm-started and only needs # N_EPOCHS_PHI. print("Solving the initial Poisson equation for E^0...") start_phi0 = time.perf_counter() key, phi_state_n = solve_poisson( key, phi_state_n, space_f, N_EPOCHS_PHI_INIT, file_name="bump_on_tail_poisson_init", retrain=False, ) end_phi0 = time.perf_counter() print(f"Solving for E^0... Done in {end_phi0 - start_phi0:.2f} seconds\n") # phi_state_star must enter the scan with the same pytree structure it will # have coming out (jax.lax.scan requires the carry's structure/dtypes to # match exactly) -- for "pinn" it already does (its own independent, # untrained net), but for "fft"/"fem" it is still the None placeholder set # above, since only phi_state_n was solved for. Since both direct solvers # ignore whatever phi_state they are handed anyway (see solve_poisson # above), reusing phi_state_n's freshly solved value is a cheap, valid seed. if phi_state_star is None: phi_state_star = phi_state_n x_diag = jnp.linspace(X_MIN, X_MAX, 512, endpoint=False)[:, None] def time_step(carry, _): """One predictor-corrector macro time step [t_n, t_n + DT]. Only the substep width DT matters for the (time-autonomous) advection fields built below, so the fixed local time_domain=(0.0, DT) that `nsl` was built with covers every step -- the absolute times are reconstructed after the scan for plotting/diagnostics. """ key, space_f, phi_state_n, phi_state_star = carry # diagnostic: electric energy of E^n, the field driving this step's prediction electric_field_n = make_electric_field(phi_state_n) e_diag = jax.vmap(electric_field_n)(x_diag) electric_energy = 0.5 * jnp.mean(e_diag**2) * (X_MAX - X_MIN) # 2. transport f^n with E^n to get the prediction f^* # `nsl` itself is never rebound here: `solve` returns a new # NeuralSemiLagrangian wrapper (holding the trained space), but the # persistent `nsl` object -- whose `.characteristic`/`.pde` are shared # by that wrapper -- must stay the same Python object across scan # iterations, so only its `.space` is extracted into a local variable. # `nsl.losses` always equals the (fixed) initial-condition fit history # followed by exactly this call's own N_EPOCHS_F entries -- since `nsl` # is never updated, that base never grows -- so the last N_EPOCHS_F # entries are this predictor substep's own loss curve. nsl.characteristic.advection_field = make_advection_field(phi_state_n) key, nsl_star = nsl.solve( key, space_f, N_EPOCHS_F, N_COLLOC_F, no_tqdm=True, **ENG_KWARGS ) space_f_star = nsl_star.space predictor_losses = nsl_star.losses.losses_history["total"][-N_EPOCHS_F:] # 3. solve Delta Psi^* = - rho[f^*] to get E^*, and set Ebar = (E^n+E^*)/2 key, phi_state_star = solve_poisson( key, phi_state_star, space_f_star, N_EPOCHS_PHI ) # 4. transport f^n (not f^*) with Ebar to get f^{n+1} nsl.characteristic.advection_field = make_averaged_advection_field( phi_state_n, phi_state_star ) key, nsl_next = nsl.solve( key, space_f, N_EPOCHS_F, N_COLLOC_F, no_tqdm=True, **ENG_KWARGS ) space_f = nsl_next.space corrector_losses = nsl_next.losses.losses_history["total"][-N_EPOCHS_F:] # 1. (for the next iteration) solve Delta Psi^{n+1} = - rho[f^{n+1}], # warm-started from E^n key, phi_state_n = solve_poisson(key, phi_state_n, space_f, N_EPOCHS_PHI) new_carry = (key, space_f, phi_state_n, phi_state_star) step_losses = jnp.concatenate([predictor_losses, corrector_losses], axis=0) return new_carry, (electric_energy, step_losses) try: from jax_tqdm import scan_tqdm time_step = scan_tqdm(NT, print_rate=1, tqdm_type="std")(time_step) except ImportError: warnings.warn( "jax_tqdm not installed, cannot display tqdm progress bar. To display a " "progress bar, install jax_tqdm." ) print(f"Running {NT} predictor-corrector steps...") start_loop = time.perf_counter() init_carry = (key, space_f, phi_state_n, phi_state_star) ( (key, space_f, phi_state_n, phi_state_star), (electric_energy_history, transport_losses_history), ) = jax.lax.scan(time_step, init_carry, jnp.arange(NT)) end_loop = time.perf_counter() print(f"Running {NT} steps... Done in {end_loop - start_loop:.2f} seconds\n") # Full loss history: the initial condition fit, followed by every predictor and # corrector transport substep's own loss curve, in chronological order (NT * 2 # substeps, each N_EPOCHS_F entries long). losses_history = jnp.concatenate( [nsl.losses.losses_history["total"], transport_losses_history.reshape((-1,))], axis=0, ) times_history = DOM_T[0] + jnp.arange(NT) * DT # Independent reference solution: classical grid-based semi-Lagrangian solver # with an FFT-based Poisson solve, used to check the NSL solution against. print("Running the reference semi-Lagrangian solver (FFT-based Poisson)...") start_ref = time.perf_counter() f_ref, x_ref, v_ref, ref_times, ref_energy = run_reference_sl(DOM_T[1]) end_ref = time.perf_counter() print( f"Running the reference semi-Lagrangian solver... " f"Done in {end_ref - start_ref:.2f} seconds\n" ) # Plot the final phase-space distribution against the reference, and the electric # energy growth, the classical diagnostic for the bump-on-tail instability. print("\nPlotting solution at final time...") plot_solution_vs_reference(space_f, f_ref, x_ref, v_ref, DOM_T[1]) plot_electric_energy(times_history, electric_energy_history, ref_times, ref_energy) plot_loss_history(losses_history) plt.show() # %%