# %% """ Vlasov-Poisson 1D — flot PINN + Poisson FFT, champ E DIFFÉRENTIABLE (bilevel). Différence avec `lagrangian_pinn_vp_fft.py` (schéma Picard figé) : - Ici E n'est PAS figé. Il est exposé comme « intermediate value » du pytree espace : ENG voit ∂E/∂θ et l'ajoute à la matrice de Gram (terme chaîne ∂resid/∂E · ∂E/∂θ). C'est l'hypergradient bilevel exact. - Le solve de Poisson `poisson_e(rho)` est enveloppé dans un `jax.custom_vjp` (théorème des fonctions implicites). Poisson étant linéaire / auto-adjoint et résolu exactement par FFT, le backward est un solve adjoint exact (obtenu ici via `jax.vjp` du corps primal — adjoint spectral sans facteur Nyquist). Mécanique d'intégration : - `VPFlowSpace.get_intermediate_values()` renvoie (E_grid,) de forme (N_T_BINS, NX_POISSON), recalculée à chaque évaluation depuis θ (jamais stop_gradient'd). La grille temporelle `_t_bins` vit dans aux_data (statique, pas un leaf → pas optimisée par ENG). - Le résidu Vlasov est POINTWISE avec model_type "x_v_t_dofsl" : le marqueur `dofsl` dit au framework que le dernier argument (E_grid) est partagé (axe vmap None). Le résidu de CI déclare aussi `dofsl` et ignore E_grid (sinon décompte d'arguments incohérent, car les intermediate values sont ajoutées à TOUS les résidus). - `check_model_type` rejette `dofsl` : on contourne en réassignant `self.model_type` APRÈS `super().__init__` (aucune re-validation à l'unflatten). """ import copy import json import os import timeit import jax import jax.numpy as jnp import matplotlib.pyplot as plt import numpy as np from scimba_jax.domains.meshless_domains.domains_1d import Segment1D from scimba_jax.linear_approximation.quad.gauss_quad import UnitSquareTensorized from scimba_jax.nonlinear_approximation.approximation_spaces.densityflowfields_approximation_spaces import ( DensityFlowFieldsApproximationSpace, ) 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.integration.monte_carlo_time import ( UniformTimeSampler, ) 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.projectors import Projector from scimba_jax.physical_models.abstract_physical_model import ( PHYSICAL_RESIDUALS_TYPE, AbstractPhysicalModel, ) from scimba_jax.physical_models.abstract_residuals import ( PARAM_FUNC_TYPE, InitialResidual, InteriorResidual, ) jax.config.update("jax_enable_x64", True) # ───────────────────────────────────────────────────────────────────────────── # Style des figures + dossier de sauvegarde des données # ───────────────────────────────────────────────────────────────────────────── # Police « LaTeX » : si une distribution LaTeX est installée, mettre USE_TEX=1 # pour un vrai rendu usetex ; sinon on utilise mathtext Computer Modern (même # allure, aucune dépendance). USE_TEX = bool(int(os.environ.get("VP_USE_TEX", "0"))) plt.rcParams.update( { "text.usetex": USE_TEX, "font.family": "serif", "font.serif": ["cmr10", "Computer Modern Roman", "DejaVu Serif"], "mathtext.fontset": "cm", "axes.formatter.use_mathtext": True, "axes.unicode_minus": False, "font.size": 12, "axes.labelsize": 13, "axes.titlesize": 13, "legend.fontsize": 11, "xtick.labelsize": 11, "ytick.labelsize": 11, "axes.linewidth": 0.8, "lines.linewidth": 1.5, "figure.dpi": 120, "savefig.dpi": 200, "savefig.bbox": "tight", } ) # Dossier où sont sauvegardées les données (npz) et les figures (pdf/png) pour # post-traitement ultérieur sans relancer l'entraînement. SAVE_DIR = os.path.join( os.path.dirname(os.path.abspath(__file__)), "results_vp_fft_bilevel" ) os.makedirs(SAVE_DIR, exist_ok=True) def _save_fig(fig, name): """Sauvegarde une figure en pdf + png dans SAVE_DIR.""" for ext in ("pdf", "png"): fig.savefig(os.path.join(SAVE_DIR, f"{name}.{ext}")) # Mode « smoke » pour valider rapidement la chaîne différentiable : # VP_SMOKE=1 python lagrangian_pinn_vp_fft_bilevel.py # → 1 flot, peu d'epochs/points, pas de référence SL ni diagnostics lourds. _SMOKE = int(os.environ.get("VP_SMOKE", "0")) # ───────────────────────────────────────────────────────────────────────────── # Paramètres physiques # ───────────────────────────────────────────────────────────────────────────── CASE = "landau_damping" if CASE == "landau_damping": K = 0.5 # nombre d'onde X_L = 2.0 * jnp.pi / K # longueur du domaine périodique ≈ 4π EPS = 0.1 # amplitude de la perturbation V_MAX = 6.0 # borne en vitesse elif CASE == "two_stream": K = 0.3 # nombre d'onde X_L = 2.0 * jnp.pi / K # longueur du domaine périodique ≈ 4π EPS = 0.05 V_MAX = 5 * jnp.pi V0 = 3.0 else: raise ValueError(f"Case {CASE!r} not implemented") def f0(xv, *_, case=CASE): """Initial distribution function f_0(x, v). case: - "landau_damping" - "two_stream" """ x = xv[..., 0] v = xv[..., 1] if case == "landau_damping": M = jnp.exp(-0.5 * v**2) / jnp.sqrt(2.0 * jnp.pi) return (1.0 + EPS * jnp.cos(K * x)) * M if case == "two_stream": M_two = 1.0 / (2.0 * jnp.sqrt(2.0 * jnp.pi)) M_two = M_two * ( jnp.exp(-((v - V0) ** 2) / 2.0) + jnp.exp(-((v + V0) ** 2) / 2.0) ) return (1.0 + EPS * jnp.cos(K * x)) * M_two raise ValueError(f"Case {case!r} not implemented") def ic_flow(x, v, *_): """Condition initiale pour le flot : identité φ(x,v,0) = (x,v). Le sampler IC passe x et v séparément — on les concatène pour obtenir (x,v). """ return jnp.concatenate([x, v], axis=-1) # ───────────────────────────────────────────────────────────────────────────── # Paramètres numériques # ───────────────────────────────────────────────────────────────────────────── N_COLLOC = 5_000 N_IC_COLLOC = 2_000 N_EPOCHS = 200 nb_models = 8 Tf = 24 if _SMOKE: N_COLLOC = 500 N_IC_COLLOC = 300 N_EPOCHS = 3 nb_models = 1 time_intervals = [ (Tf * i / nb_models, Tf * (i + 1) / nb_models) for i in range(nb_models) ] dim_x = 1 dim_v = 1 params_dim = 0 # ───────────────────────────────────────────────────────────────────────────── # Grille Poisson FFT # ───────────────────────────────────────────────────────────────────────────── NX_POISSON = 32 NQ_POISSON = 128 N_T_BINS = 10 x_poisson = jnp.linspace(0.0, float(X_L), NX_POISSON, endpoint=False) dx_poisson = float(X_L) / NX_POISSON kx_rfft = jnp.fft.rfftfreq(NX_POISSON) * NX_POISSON * 2.0 * jnp.pi / float(X_L) kx_safe = jnp.where(kx_rfft != 0.0, kx_rfft, 1.0) _quad = UnitSquareTensorized(dim=1, order=NQ_POISSON) quad_v_pts = (-V_MAX + 2.0 * V_MAX * _quad.volumic_points).ravel() quad_v_w = (2.0 * V_MAX * _quad.volumic_weights).ravel() _xx_xv = jnp.repeat(x_poisson, NQ_POISSON) _vv_xv = jnp.tile(quad_v_pts, NX_POISSON) # ───────────────────────────────────────────────────────────────────────────── # Poisson FFT — version classique (diagnostics / inférence) # ───────────────────────────────────────────────────────────────────────────── def solve_poisson_fft(rho: jnp.ndarray): rho_hat = jnp.fft.rfft(rho - 1.0) phi_hat = jnp.where(kx_rfft != 0.0, rho_hat / kx_safe**2, 0.0 + 0.0j) E_hat = jnp.where(kx_rfft != 0.0, -1j * kx_safe * phi_hat, 0.0 + 0.0j) return jnp.fft.irfft(phi_hat, n=NX_POISSON), jnp.fft.irfft(E_hat, n=NX_POISSON) # ───────────────────────────────────────────────────────────────────────────── # Poisson FFT DIFFÉRENTIABLE via custom_vjp (théorème des fonctions implicites) # ───────────────────────────────────────────────────────────────────────────── # Niveau bas (inner problem) : φ solution de -Δφ = ρ - 1, puis E = -∂_x φ. # En Fourier (rfft, longueur NX_POISSON) l'opérateur est diagonal : # φ̂ = ρ̂ / k² (mode k=0 mis à zéro) # Ê = -i k φ̂ = -i ρ̂ / k # L'application ρ ↦ E est linéaire. L'IFT donne alors # dφ/dρ = (-Δ)⁻¹ (auto-adjoint) # donc le VJP (backward) est le MÊME solve appliqué au cotangent : c'est un # solve adjoint de Poisson. On l'obtient ici de façon exacte via `jax.vjp` du # corps primal (évite tout facteur Nyquist d'un adjoint spectral écrit à la # main). Le wrapper custom_vjp documente la structure et généralise au cas # non-linéaire (où le backward redeviendrait un vrai solve adjoint). def _poisson_e_primal(rho: jnp.ndarray) -> jnp.ndarray: rho_hat = jnp.fft.rfft(rho - 1.0) phi_hat = jnp.where(kx_rfft != 0.0, rho_hat / kx_safe**2, 0.0 + 0.0j) # φ̂ E_hat = jnp.where(kx_rfft != 0.0, -1j * kx_safe * phi_hat, 0.0 + 0.0j) # -ik φ̂ return jnp.fft.irfft(E_hat, n=NX_POISSON) @jax.custom_vjp def poisson_e(rho: jnp.ndarray) -> jnp.ndarray: """E(ρ) = -∂_x (-Δ)⁻¹ (ρ-1), différentiable (backward = solve adjoint).""" return _poisson_e_primal(rho) def _poisson_e_fwd(rho: jnp.ndarray): return _poisson_e_primal(rho), rho def _poisson_e_bwd(rho: jnp.ndarray, g: jnp.ndarray): # IFT / adjoint : ρ̄ = L^* g, avec L : ρ ↦ E linéaire. Pour un opérateur # linéaire, le VJP autodiff du corps primal EST l'adjoint exact. _, vjp = jax.vjp(_poisson_e_primal, rho) return vjp(g) poisson_e.defvjp(_poisson_e_fwd, _poisson_e_bwd) # ───────────────────────────────────────────────────────────────────────────── # Interpolation Catmull-Rom # ───────────────────────────────────────────────────────────────────────────── def catmull_rom_interp(field_grid: jnp.ndarray, x_query: jnp.ndarray) -> jnp.ndarray: Nx = field_grid.shape[0] dx = float(X_L) / Nx pos = (x_query % float(X_L)) / dx i1 = jnp.floor(pos).astype(jnp.int32) t = pos - i1 i0 = (i1 - 1) % Nx i2 = (i1 + 1) % Nx i3 = (i1 + 2) % Nx f0_ = field_grid[i0] f1_ = field_grid[i1] f2_ = field_grid[i2] f3_ = field_grid[i3] return 0.5 * ( (-f0_ + 3.0 * f1_ - 3.0 * f2_ + f3_) * t**3 + (2.0 * f0_ - 5.0 * f1_ + 4.0 * f2_ - f3_) * t**2 + (-f0_ + f2_) * t + 2.0 * f1_ ) # ───────────────────────────────────────────────────────────────────────────── # compute_rho — composition des flots 0..i (φ_i interne, précédents gelés) # ───────────────────────────────────────────────────────────────────────────── def compute_rho(space, t_scalar) -> jnp.ndarray: t_val = jnp.asarray(t_scalar).reshape(()) tt = jnp.broadcast_to(t_val, (_xx_xv.shape[0],)) xv_grid = jnp.stack([_xx_xv, _vv_xv], axis=-1) # (NX*NQ, 2) idx = space.idx_current_flow def density_at_point(xv_pt, t_pt): # Temps LOCAL t - t_start par flot (cohérent avec la normalisation du # parent densityflowfields ; requis dès qu'on utilise un sympnet). t_loc_i = t_pt.reshape(1) - space.time_intervals[idx, 0:1] extra_i = jnp.concatenate([t_loc_i, jnp.zeros((0,))], axis=-1) result = space._apply_model(space.models[idx], xv_pt, extra_i) det = space._abs_det_jac_wrt_x(space.models[idx], xv_pt, extra_i) for k in range(idx - 1, -1, -1): t_k = space.time_intervals[k, 1:2] - space.time_intervals[k, 0:1] extra_k = jnp.concatenate([t_k, jnp.zeros((0,))], axis=-1) det = det * jax.lax.stop_gradient( space._abs_det_jac_wrt_x(space.models[k], result, extra_k) ) result = jax.lax.stop_gradient( space._apply_model(space.models[k], result, extra_k) ) return space.initial_density(result, t_pt.reshape(1), jnp.zeros((0,))) * det f_vals = jax.vmap(jax.checkpoint(density_at_point))(xv_grid, tt) return (f_vals.reshape(NX_POISSON, NQ_POISSON) * quad_v_w[None, :]).sum(axis=-1) def compute_e_grid(space, t_scalar) -> jnp.ndarray: """E(t) non différentiable (diagnostics) — via solve_poisson_fft.""" _, E = solve_poisson_fft(compute_rho(space, t_scalar)) return E def compute_e_grid_diff(space, t_scalar) -> jnp.ndarray: """E(t) DIFFÉRENTIABLE w.r.t. θ — via poisson_e (custom_vjp IFT).""" return poisson_e(compute_rho(space, t_scalar)) def compute_rho_inference(space, t_scalar) -> jnp.ndarray: xv_grid = jnp.stack([_xx_xv, _vv_xv], axis=-1) t_val = jnp.asarray(t_scalar).reshape(()) tt = jnp.broadcast_to(t_val, (xv_grid.shape[0], 1)) xvt = jnp.concatenate([xv_grid, tt], axis=-1) f_vals = jax.vmap(lambda a: space.compute_density_inference(space, a))(xvt) return (f_vals.reshape(NX_POISSON, NQ_POISSON) * quad_v_w[None, :]).sum(axis=-1) def compute_electric_energy(rho, dx): Nx = rho.shape[0] kx = jnp.fft.fftfreq(Nx) * Nx * 2.0 * jnp.pi / float(X_L) kx_s = jnp.where(kx != 0, kx, 1.0) rho_hat = jnp.fft.fft(rho - 1.0) E_hat = jnp.where(kx != 0, -1j * rho_hat / kx_s, 0.0 + 0.0j) return 0.5 * jnp.sum(jnp.fft.ifft(E_hat).real ** 2) * dx # ───────────────────────────────────────────────────────────────────────────── # VPFlowSpace — E exposé comme intermediate value différentiable # ───────────────────────────────────────────────────────────────────────────── class VPFlowSpace(DensityFlowFieldsApproximationSpace): """ Étend DensityFlowFieldsApproximationSpace avec : 1. get_intermediate_values : renvoie la grille E(θ) DIFFÉRENTIABLE (forme (N_T_BINS, NX_POISSON)) → ENG voit ∂E/∂θ. 2. compute_density_training : compose les flots 0..i (φ_i le plus interne). La grille temporelle `_t_bins` est dans aux_data (statique) : ce n'est PAS un leaf, donc ENG ne l'optimise pas. Elle est fixée par flot avant `project`. """ def __init__(self, *args, n_t_bins: int = N_T_BINS, **kwargs): super().__init__(*args, **kwargs) self._n_t_bins = n_t_bins self._t_bins = jnp.zeros((n_t_bins,)) # placeholder, fixé par flot # ── Pytree ─────────────────────────────────────────────────────────────── # Plus de tree_flatten/tree_unflatten ici. `_t_bins` etait mis en aux_data # pour que ENG ne l'optimise pas ; ce n'est plus l'emplacement qui le dit # mais le ROLE : la classification generique en fait un child, et # `auto_partition` le laisse gele parce qu'il n'est pas declare # `trainable`. `_n_t_bins`, un int, part en aux_data tout seul. # ── inference_mode ──────────────────────────────────────────────────────── # compute_density est une propriete cote bibliotheque (une methode liee # rangee en aux_data se compare par identite et coute une recompilation a # chaque aller-retour pytree) : on la SURCHARGE au lieu de l'affecter. @property def compute_density(self): if self._inference_mode: return self.compute_density_inference return self.compute_density_training @property def inference_mode(self) -> bool: return self._inference_mode @inference_mode.setter def inference_mode(self, value: bool): self._inference_mode = bool(value) # ── compute_density en mode entraînement ───────────────────────────────── def compute_density_training(self, space, *args): cargs = jnp.concatenate(args, axis=-1) x = space.get("x", cargs) v = space.get("v", cargs) t = space.get("t", cargs) mu = space.get("mu", cargs) result = jnp.concatenate([x, v], axis=-1) det = jnp.array(1.0) for k in range(space.idx_current_flow, -1, -1): t_k = t if k == space.idx_current_flow else space.time_intervals[k, 1:2] t_k = ( t_k - space.time_intervals[k, 0:1] ) # temps local (cf. parent normalisé) extra = jnp.concatenate([t_k, mu], axis=-1) det = det * space._abs_det_jac_wrt_x(space.models[k], result, extra) result = space._apply_model(space.models[k], result, extra) return space.initial_density(result, t, mu) * det # ── Intermediate values : la grille E(θ) différentiable ─────────────────── def get_intermediate_values(self): E_grid = jax.lax.map(lambda t: compute_e_grid_diff(self, t), self._t_bins) return (E_grid,) # (N_T_BINS, NX_POISSON) def get_intermediate_values_shapes(self): return ((self._n_t_bins, NX_POISSON),) # ───────────────────────────────────────────────────────────────────────────── # Tests unitaires # ───────────────────────────────────────────────────────────────────────────── def test_poisson_fft(): print("\n" + "=" * 60) print("=== Step 1: solve_poisson_fft + poisson_e (custom_vjp) ===") rho_test = 1.0 + EPS * jnp.cos(K * x_poisson) phi_exact = (EPS / K**2) * jnp.cos(K * x_poisson) phi_exact = phi_exact - phi_exact.mean() E_exact = (EPS / K) * jnp.sin(K * x_poisson) phi_num, E_num = solve_poisson_fft(rho_test) phi_num = phi_num - phi_num.mean() err_phi = float(jnp.abs(phi_num - phi_exact).max()) err_E = float(jnp.abs(E_num - E_exact).max()) print( f" max|phi - phi_exact| = {err_phi:.2e} {'OK' if err_phi < 1e-8 else 'FAIL'}" ) print(f" max|E - E_exact | = {err_E:.2e} {'OK' if err_E < 1e-8 else 'FAIL'}") # poisson_e doit coïncider avec solve_poisson_fft sur le forward err_fwd = float(jnp.abs(poisson_e(rho_test) - E_num).max()) print( f" max|poisson_e - E_fft| = {err_fwd:.2e} " f"{'OK' if err_fwd < 1e-12 else 'FAIL'}" ) # backward custom_vjp vs vjp direct du corps primal + FD def loss_e(rho): return jnp.sum(poisson_e(rho) ** 2) grad_custom = jax.grad(loss_e)(rho_test) def loss_e_primal(rho): return jnp.sum(_poisson_e_primal(rho) ** 2) grad_primal = jax.grad(loss_e_primal)(rho_test) err_adj = float(jnp.abs(grad_custom - grad_primal).max()) print( f" max|grad_customvjp - grad_autodiff| = {err_adj:.2e} " f"{'OK' if err_adj < 1e-10 else 'FAIL'}" ) eps_fd = 1e-5 grad_fd = jnp.array( [ (loss_e(rho_test.at[j].add(eps_fd)) - loss_e(rho_test.at[j].add(-eps_fd))) / (2 * eps_fd) for j in range(NX_POISSON) ] ) err_fd = float(jnp.abs(grad_custom - grad_fd).max()) print( f" max|grad_customvjp - grad_FD| = {err_fd:.2e} " f"{'OK' if err_fd < 1e-5 else 'FAIL'}" ) def test_catmull_rom(): print("\n" + "=" * 60) print("=== Step 2: Catmull-Rom interpolation ===") E_grid = (EPS / K) * jnp.sin(K * x_poisson) err_nodes = float(jnp.abs(catmull_rom_interp(E_grid, x_poisson) - E_grid).max()) print( f" Erreur sur les noeuds : {err_nodes:.2e} {'OK' if err_nodes < 1e-10 else 'FAIL'}" ) x_fine = jnp.linspace(0.0, float(X_L), 4 * NX_POISSON, endpoint=False) E_fine = (EPS / K) * jnp.sin(K * x_fine) err_fine = float(jnp.abs(catmull_rom_interp(E_grid, x_fine) - E_fine).max()) h = float(X_L) / NX_POISSON print( f" Erreur sur grille fine : {err_fine:.2e} (h^4={h**4:.2e}) " f"{'OK' if err_fine < 10 * h**4 else 'FAIL'}" ) def test_compute_rho(space): print("\n" + "=" * 60) print("=== Step 3: compute_rho ===") space.idx_current_flow = 0 space.inference_mode = False rho_num = compute_rho(space, 0.0) rho_exact = 1.0 + EPS * jnp.cos(K * x_poisson) err = float(jnp.abs(rho_num - rho_exact).max()) print( f" max|rho(t=0) - rho_exact| = {err:.2e} " f"(untrained network {'OK' if err < 0.1 else 'FAIL'})" ) mass = float(rho_num.sum() * dx_poisson) print( f" integral rho dx = {mass:.6f} (expected {float(X_L):.6f}) " f"{'OK' if abs(mass - float(X_L)) < 0.1 else 'FAIL'}" ) def test_compute_e(space): print("\n" + "=" * 60) print("=== Step 4: compute_E (diff vs non-diff) ===") space.idx_current_flow = 0 space.inference_mode = False E_num = compute_e_grid(space, 0.0) E_diff = compute_e_grid_diff(space, 0.0) err = float(jnp.abs(E_num - E_diff).max()) print( f" max|compute_e_grid - compute_e_grid_diff| = {err:.2e} " f"{'OK' if err < 1e-12 else 'FAIL'}" ) def loss_e(s): return jnp.sum(compute_e_grid_diff(s, 0.0) ** 2) grad = jax.grad(loss_e)(space) leaves = jax.tree_util.tree_leaves(grad) has_nan = any( bool(jnp.any(jnp.isnan(ell))) for ell in leaves if hasattr(ell, "shape") ) max_grad = max( ( float(jnp.abs(ell).max()) for ell in leaves if hasattr(ell, "shape") and ell.size > 1 ), default=0.0, ) print( f" max||grad_theta E^2|| = {max_grad:.2e} " f"{'NAN' if has_nan else ('OK (E flows in the gradient)' if max_grad > 1e-10 else 'CUT')}" ) def test_intermediate_values(space): print("\n" + "=" * 60) print("=== Step 5: intermediate values E(theta) + dE/dtheta ===") space.idx_current_flow = 0 space.inference_mode = False space._t_bins = jnp.linspace(*time_intervals[0], N_T_BINS) (E_grid,) = space.get_intermediate_values() has_nan = bool(jnp.any(jnp.isnan(E_grid))) print( f" E_grid shape : {E_grid.shape} (attendu ({N_T_BINS}, {NX_POISSON})) " f"{'OK' if E_grid.shape == (N_T_BINS, NX_POISSON) else 'FAIL'}" ) print(f" NaN in E_grid : {'YES' if has_nan else 'NO'}") print( f" shapes() = {space.get_intermediate_values_shapes()} " f"max|E_grid| = {float(jnp.abs(E_grid).max()):.4e}" ) # ∂E_grid/∂θ via le mécanisme du framework (utilisé par ENG) jac = space.jac_intermediate_value_theta(0) n_out = N_T_BINS * NX_POISSON ok_shape = jac.shape[0] == n_out max_jac = float(jnp.abs(jac).max()) has_nan_j = bool(jnp.any(jnp.isnan(jac))) print( f" jac_intermediate_value_theta(0) shape : {jac.shape} " f"(lignes attendues {n_out}) {'OK' if ok_shape else 'FAIL'}" ) print( f" max|dE/dtheta| = {max_jac:.2e} " f"{'NAN' if has_nan_j else ('OK (E depends on theta)' if max_jac > 1e-10 else 'CUT')}" ) def _ref_rho_full_compose(space, t_scalar): """ρ de référence : pleine composition via compute_density_training (φ_i interne).""" tt = jnp.full((_xx_xv.shape[0], 1), float(t_scalar)) args = jnp.concatenate([_xx_xv[:, None], _vv_xv[:, None], tt], axis=-1) f = jax.vmap(lambda a: space.compute_density_training(space, a))(args) return (f.reshape(NX_POISSON, NQ_POISSON) * quad_v_w[None, :]).sum(axis=-1) def test_compose_order(space): """Vérifie que compute_rho respecte l'ordre de composition (φ_i interne) pour i≥1.""" print("\n" + "=" * 60) print("=== Step 6: compute_rho composition order (idx=1) ===") if len(time_intervals) < 2: print(" (single flow - test skipped)") return space.idx_current_flow = 1 space.inference_mode = False t_test = float(time_intervals[1][0]) rho_fast = compute_rho(space, t_test) rho_ref = _ref_rho_full_compose(space, t_test) err = float(jnp.abs(rho_fast - rho_ref).max()) print( f" max|compute_rho - full-compose ref.| = {err:.2e} " f"{'OK' if err < 1e-9 else 'FAIL'}" ) space.idx_current_flow = 0 # ───────────────────────────────────────────────────────────────────────────── # Résidu Vlasov — POINTWISE avec E_grid comme intermediate value (dofsl) # ───────────────────────────────────────────────────────────────────────────── def _interp_e_at(E_grid, x, t, t0_, t1_): """Interpole E au point (x, t) depuis la grille (N_T_BINS, NX_POISSON).""" n_t_bins = E_grid.shape[0] dt_bin = (t1_ - t0_) / (n_t_bins - 1) idx_f = (t[0] - t0_) / dt_bin i0 = jnp.clip(jnp.floor(idx_f).astype(jnp.int32), 0, n_t_bins - 2) a = idx_f - i0 E_row = (1.0 - a) * E_grid[i0] + a * E_grid[i0 + 1] # (NX_POISSON,) return catmull_rom_interp(E_row, x.reshape(1))[0] # scalaire class VlasovResidualFFT(InteriorResidual): def __init__(self, domain, time_domain: tuple[float, float], f_rhs=None): # model_type "x_v_t" pour passer check_model_type, puis on ajoute `dofsl` # (marqueur de l'intermediate value E_grid, axe vmap None) — aucune # re-validation n'a lieu ensuite (ni à l'unflatten). super().__init__( domain=domain, size=2, model_type="x_v_t", f_rhs=f_rhs, time_domain=time_domain, ) self.model_type = "x_v_t_dofsl" self._t0 = float(time_domain[0]) self._t1 = float(time_domain[1]) def construct_residual(self, *args) -> PARAM_FUNC_TYPE: inv_phi = args[0] phi_x, phi_v = inv_phi.components() dims = dict(phi_x.dims) dims["dofsl"] = 1 # requis par la validation dims/f_type dt_px_op = phi_x.d_t() dt_pv_op = phi_v.d_t() dx_px_op = phi_x.gradient("x") dx_pv_op = phi_v.gradient("x") dv_px_op = phi_x.gradient("v") dv_pv_op = phi_v.gradient("v") t0_, t1_ = self._t0, self._t1 def pointwise_fn(space, x, v, t, E_grid): # x, v, t : (1,) ; E_grid : (N_T_BINS, NX_POISSON) partagé E_val = _interp_e_at(E_grid, x, t, t0_, t1_) v0 = v[0] res_x = ( dt_px_op(space, x, v, t) + v0 * dx_px_op(space, x, v, t)[0] + E_val * dv_px_op(space, x, v, t)[0] ) res_v = ( dt_pv_op(space, x, v, t) + v0 * dx_pv_op(space, x, v, t)[0] + E_val * dv_pv_op(space, x, v, t)[0] ) return jnp.stack([res_x, res_v]) # (2,) return ParamVecFunction( size=2, dim=dims, fn=pointwise_fn, f_type="x_v_t_dofsl", ) # ───────────────────────────────────────────────────────────────────────────── # CI du flot — déclare aussi `dofsl` et ignore E_grid # ───────────────────────────────────────────────────────────────────────────── class ICFlowResidual(InitialResidual): def __init__(self, domain, time_domain=(0.0,), f_rhs=None): super().__init__( domain=domain, time_domain=( (time_domain,) if isinstance(time_domain, float) else time_domain ), size=2, model_type="x_v_t", f_rhs=f_rhs, ) # super() a réduit à "x_v" (sans temps) ; on ajoute `dofsl` car les # intermediate values sont ajoutées aux args de TOUS les résidus. self.model_type = "x_v_dofsl" def construct_residual(self, *args) -> PARAM_FUNC_TYPE: r = args[0].set_t_0(self.time_domain[0]) # ParamVecFunction f_type "x_v" dims = dict(r.dims) dims["dofsl"] = 1 def ic_fn(space, x, v, E_grid): # E_grid ignoré return r(space, x, v) return ParamVecFunction(size=2, dim=dims, fn=ic_fn, f_type="x_v_dofsl") # ───────────────────────────────────────────────────────────────────────────── # Modèle physique # ───────────────────────────────────────────────────────────────────────────── class VlasovPoissonModelFFT(AbstractPhysicalModel): def __init__(self, main_domain, time_domain, f_ic_rhs=None): super().__init__(main_domain=main_domain, time_domain=time_domain) self._main_domain = main_domain self._f_ic_rhs = f_ic_rhs self._build_residuals(time_domain) def _build_residuals(self, time_domain): label = self._main_domain.get_label() self.physical_residuals: PHYSICAL_RESIDUALS_TYPE = { label: VlasovResidualFFT( domain=self._main_domain, time_domain=time_domain, ), "ic " + label: ICFlowResidual( domain=self._main_domain, time_domain=(time_domain[0],), f_rhs=self._f_ic_rhs, ), } def renew_time_domain(self, new_time_domain): self.time_domain = new_time_domain self._build_residuals(new_time_domain) # ───────────────────────────────────────────────────────────────────────────── # Domaines + quadrature + réseaux + espace # ───────────────────────────────────────────────────────────────────────────── key = jax.random.PRNGKey(42) domain_x = Segment1D(jnp.array([[0.0, float(X_L)]]), is_main_domain=True) domain_v = Segment1D(jnp.array([[-V_MAX, V_MAX]])) N_QUAD_V = 96 quad_velocity = UnitSquareTensorized(dim=dim_v, order=N_QUAD_V) quad_velocity.volumic_points = -V_MAX + 2.0 * V_MAX * quad_velocity.volumic_points quad_velocity.volumic_weights = 2.0 * V_MAX * quad_velocity.volumic_weights keys = jax.random.split(key, nb_models + 1) # Config des MLP, factorisée pour pouvoir reconstruire des squelettes identiques # au rechargement (tree_deserialise_leaves a besoin d'un modèle « modèle »). MLP_CONFIG = dict( in_size=dim_x + dim_v + 1, hidden_sizes=[30, 30], out_size=dim_x + dim_v, activation="tanh", embedding="periodic", periods=(float(X_L),), embedding_axes=(0,), ) def build_flow_models(keys_): """Construit la liste des MLP de flot (squelettes ou modèles à entraîner).""" return [MLP(key=keys_[i], **MLP_CONFIG) for i in range(nb_models)] flow_models = build_flow_models(keys) space = VPFlowSpace( dim=dim_x, velocity_dim=dim_v, params_dim=params_dim, models=flow_models, list_models_fields=[], add_id=True, time_intervals=time_intervals, time_continuous=True, parametric_dependence=False, initial_density=f0, moment_computation=True, quadrature_velocity=quad_velocity, scan_moments=False, n_t_bins=N_T_BINS, ) # ───────────────────────────────────────────────────────────────────────────── # Tests # ───────────────────────────────────────────────────────────────────────────── # %% print("\n" + "=" * 60) print("=== TESTS STEPS 1-6 ===") print("=" * 60) test_poisson_fft() test_catmull_rom() space.idx_current_flow = 0 space.inference_mode = False test_compute_rho(space) test_compute_e(space) test_intermediate_values(space) test_compose_order(space) print("\n=== Tests OK - moving to training (differentiable E) ===\n") # ───────────────────────────────────────────────────────────────────────────── # Entraînement # ───────────────────────────────────────────────────────────────────────────── initial_t_domain = (0.0, time_intervals[0][1]) sampler = TensorizedSampler( [ DomainSampler(domain_x), UniformVelocitySamplerOnCuboid(domain_v), UniformTimeSampler(initial_t_domain), ], bc=False, ic=True, model_type="x_v_t", ) model = VlasovPoissonModelFFT( main_domain=domain_x, time_domain=initial_t_domain, f_ic_rhs=ic_flow, ) weights = { "interior": [2.0, 2.0], "ic interior": [5.0, 5.0], } list_of_pinns = [] pinn_total_time = 0.0 for i in range(nb_models): space.idx_current_flow = i domain_t = time_intervals[i] print(f"\n=== Flow {i + 1}/{nb_models} - t in {domain_t} ===") sampler.renew_sampler(UniformTimeSampler(domain_t), "t") model.renew_time_domain(domain_t) space.inference_mode = False # Grille temporelle du flot courant (statique, lue par get_intermediate_values) space._t_bins = jnp.linspace(float(domain_t[0]), float(domain_t[1]), N_T_BINS) pinn = Projector( model, space, sampler, optimizer="ENG", weights=weights, one_loss_per_residual=True, matrix_regularization=1.0e-6, ) start = timeit.default_timer() key, pinn = pinn.project( key, space, n_epochs=N_EPOCHS, n_colloc=N_COLLOC, n_ic_colloc=N_IC_COLLOC ) elapsed = timeit.default_timer() - start # Propage les poids entraînés : le flot suivant compose avec les flots # précédents ENTRAÎNÉS et gelés. space = pinn.space space._t_bins = jnp.linspace(float(domain_t[0]), float(domain_t[1]), N_T_BINS) list_of_pinns.append(copy.deepcopy(pinn)) pinn_total_time += elapsed print(f" time : {elapsed:.1f}s") print(f"\n=== PINN FFT bilevel: total time: {pinn_total_time:.1f}s ===") # ───────────────────────────────────────────────────────────────────────────── # Sauvegarde des flots entraînés (pour ré-évaluation à n'importe quel temps) # ───────────────────────────────────────────────────────────────────────────── # `space` (= dernier pinn.space) contient les nb_models flots entraînés et gelés # dans `space.models`. On sérialise leurs feuilles (.npz) + les métadonnées qui # permettent de reconstruire des squelettes identiques au rechargement. FLOWS_PATH = os.path.join(SAVE_DIR, "flows.npz") META_PATH = os.path.join(SAVE_DIR, "flows_meta.json") def tree_serialise_leaves(path, tree): """Sauvegarde les feuilles d'un pytree dans un .npz (ordre du treedef).""" leaves = jax.tree_util.tree_leaves(tree) np.savez(path, *[np.asarray(leaf) for leaf in leaves]) def tree_deserialise_leaves(path, skeleton): """Recharge les feuilles depuis un .npz dans la structure de `skeleton`.""" data = np.load(path) names = sorted(data.files, key=lambda s: int(s.split("_")[1])) leaves = [jnp.asarray(data[name]) for name in names] treedef = jax.tree_util.tree_structure(skeleton) return jax.tree_util.tree_unflatten(treedef, leaves) tree_serialise_leaves(FLOWS_PATH, space.models) with open(META_PATH, "w") as fh: json.dump( { "case": CASE, "K": float(K), "X_L": float(X_L), "EPS": float(EPS), "V_MAX": float(V_MAX), "dim_x": dim_x, "dim_v": dim_v, "params_dim": params_dim, "nb_models": nb_models, "Tf": float(Tf), "time_intervals": [[float(a), float(b)] for a, b in time_intervals], "n_t_bins": N_T_BINS, "mlp_config": { **MLP_CONFIG, "periods": [float(p) for p in MLP_CONFIG["periods"]], }, }, fh, indent=2, ) print(f"Trained flows saved in {FLOWS_PATH}") def load_flows(save_dir=SAVE_DIR): """Recharge les flots entraînés et renvoie un VPFlowSpace prêt pour l'inférence. Usage (dans une session où ce module est importé, training gardé sous main) : sp = load_flows() sp.idx_current_flow = sp.nb_models - 1 # tout l'horizon composé rho = compute_rho_inference(sp, t) # à n'importe quel temps t ∈ [0, Tf] Reconstruit des squelettes via `build_flow_models`, y recharge les feuilles sérialisées, puis rebâtit l'espace avec la configuration LUE DANS LE JSON (intervalles de vie des flots, nb_models, n_t_bins…) — pas les globales du module, pour rester correct même dans une session découplée. """ with open(os.path.join(save_dir, "flows_meta.json")) as fh: meta = json.load(fh) saved_time_intervals = [tuple(ab) for ab in meta["time_intervals"]] skeleton = build_flow_models(jax.random.split(jax.random.PRNGKey(0), nb_models + 1)) trained = tree_deserialise_leaves(os.path.join(save_dir, "flows.npz"), skeleton) sp = VPFlowSpace( dim=dim_x, velocity_dim=dim_v, params_dim=params_dim, models=trained, list_models_fields=[], add_id=True, time_intervals=saved_time_intervals, time_continuous=True, parametric_dependence=False, initial_density=f0, moment_computation=True, quadrature_velocity=quad_velocity, scan_moments=False, n_t_bins=meta["n_t_bins"], ) sp.inference_mode = True return sp # ───────────────────────────────────────────────────────────────────────────── # SL référence — calculé UNE SEULE FOIS # ───────────────────────────────────────────────────────────────────────────── NX_SL = 128 NV_SL = 128 DT_SL = 0.05 def run_sl_vlasov_poisson(t_checkpoints, Nx=NX_SL, Nv=NV_SL, dt=DT_SL): x_grid = jnp.linspace(0.0, float(X_L), Nx, endpoint=False) v_grid = jnp.linspace(-V_MAX, V_MAX, Nv, endpoint=False) dv_sl = 2.0 * V_MAX / Nv xx, vv = jnp.meshgrid(x_grid, v_grid, indexing="ij") if CASE == "landau_damping": f = ( (1.0 + EPS * jnp.cos(K * xx)) * jnp.exp(-0.5 * vv**2) / jnp.sqrt(2.0 * jnp.pi) ) elif CASE == "two_stream": f = ( (1.0 + EPS * jnp.cos(K * xx)) * (jnp.exp(-0.5 * (vv - V0) ** 2) + jnp.exp(-0.5 * (vv + V0) ** 2)) / (2 * jnp.sqrt(2.0 * jnp.pi)) ) else: raise ValueError(f"Case SL {CASE!r} not implemented") kx = jnp.fft.fftfreq(Nx) * Nx * 2.0 * jnp.pi / float(X_L) kv = jnp.fft.fftfreq(Nv) * Nv * 2.0 * jnp.pi / (2.0 * V_MAX) kx_s = jnp.where(kx != 0, kx, 1.0) def compute_phi_e(f): rho = jnp.sum(f, axis=1) * dv_sl rho_hat = jnp.fft.fft(rho - 1.0) phi_hat = jnp.where(kx != 0, rho_hat / kx_s**2, 0.0 + 0.0j) E_hat = jnp.where(kx != 0, -1j * kx_s * phi_hat, 0.0 + 0.0j) return rho, jnp.fft.ifft(phi_hat).real, jnp.fft.ifft(E_hat).real def advect_x(f, dt_sub): return jnp.fft.ifft( jnp.fft.fft(f, axis=0) * jnp.exp(-1j * kx[:, None] * v_grid[None, :] * dt_sub), axis=0, ).real def advect_v(f, E, dt_sub): return jnp.fft.ifft( jnp.fft.fft(f, axis=1) * jnp.exp(-1j * kv[None, :] * E[:, None] * dt_sub), axis=1, ).real @jax.jit def step(f): f = advect_x(f, dt / 2.0) _, _, E = compute_phi_e(f) f = advect_v(f, E, dt) f = advect_x(f, dt / 2.0) rho, phi, _ = compute_phi_e(f) return f, rho, phi checkpoint_steps = [int(t / dt + 0.5) for t in t_checkpoints] t_actual = [s * dt for s in checkpoint_steps] n_steps = max(checkpoint_steps) print(f"SL ({Nx}x{Nv}, dt={dt}): {n_steps} steps") step_to_idxs: dict[int, list[int]] = {} for idx, s in enumerate(checkpoint_steps): step_to_idxs.setdefault(s, []).append(idx) results: list = [None] * len(t_checkpoints) rho0, phi0, _ = compute_phi_e(f) if 0 in step_to_idxs: for idx in step_to_idxs[0]: results[idx] = (f, rho0, phi0) for i in range(n_steps): f, rho, phi = step(f) if i + 1 in step_to_idxs: for idx in step_to_idxs[i + 1]: results[idx] = (f, rho, phi) if (i + 1) % max(1, n_steps // 5) == 0: print(f" t = {(i + 1) * dt:.2f}") return results, t_actual, x_grid, v_grid # ───────────────────────────────────────────────────────────────────────────── # Diagnostics # ───────────────────────────────────────────────────────────────────────────── # %% if not _SMOKE: pinn_last = list_of_pinns[-1] sp = pinn_last.space sp.idx_current_flow = len(list_of_pinns) - 1 sp.inference_mode = True # ── Loss curves ────────────────────────────────────────────────────────── loss_histories = {} for i, pinn in enumerate(list_of_pinns): pinn.space.idx_current_flow = i fig, ax = plt.subplots(figsize=(7, 3)) for label, hist in pinn.losses.losses_history.items(): ax.semilogy(hist, label=label) loss_histories[f"flow{i}_{label}"] = np.asarray(hist) ax.set_xlabel("epoch") ax.set_ylabel("loss") ax.legend() ax.grid(True, alpha=0.3) plt.tight_layout() _save_fig(fig, f"loss_flow{i}") plt.show() np.savez(os.path.join(SAVE_DIR, "loss_histories.npz"), **loss_histories) # ── SL calculé une seule fois ───────────────────────────────────────────── n_t = 80 t_diag = [0.0, Tf / 4, Tf / 2, 3 * Tf / 4, Tf] t_energy = list(np.linspace(0.0, float(Tf), n_t)) t_combined = t_energy + t_diag print("\nComputing SL (once)...") sl_combined, t_actual_combined, x_sl, v_sl = run_sl_vlasov_poisson(t_combined) x_sl = np.array(x_sl) v_sl = np.array(v_sl) sl_energy = sl_combined[:n_t] sl_diag = sl_combined[n_t:] t_actual_energy = t_actual_combined[:n_t] t_actual_diag = t_actual_combined[n_t:] # ── Énergie électrique ──────────────────────────────────────────────────── dv_sl_ = 2.0 * V_MAX / NV_SL dx_sl_ = float(X_L) / NX_SL E_sl_energy = [ float(compute_electric_energy(jnp.array(f).sum(axis=1) * dv_sl_, dx_sl_)) for f, _, _ in sl_energy ] print("PINN FFT energy...") @jax.jit def compute_all_energies(s, t_arr): def energy_at_t(t): rho = compute_rho_inference(s, t) return compute_electric_energy(rho, dx_poisson) return jax.lax.map(energy_at_t, t_arr) t_arr_jnp = jnp.array(t_energy) E_pinn_energy = list(np.array(compute_all_energies(sp, t_arr_jnp))) gamma_th = -0.1533 t_arr = np.array(t_energy) E0_ref = max(E_sl_energy[0], 1e-10) fig_e, ax_e = plt.subplots(figsize=(9, 4)) ax_e.semilogy(np.array(t_actual_energy), E_sl_energy, "steelblue", lw=2, label="SL") ax_e.semilogy(t_arr, E_pinn_energy, "tomato", lw=1.5, label="PINN-FFT") ax_e.semilogy( t_arr, E0_ref * np.exp(2 * gamma_th * t_arr), "k--", lw=1, label=rf"$\gamma = {gamma_th}$", ) ax_e.set_xlabel(r"$t$") ax_e.set_ylabel(r"$\frac{1}{2}\int E^2\,\mathrm{d}x$") ax_e.legend() ax_e.grid(True, which="both", alpha=0.2) plt.tight_layout() _save_fig(fig_e, "electric_energy") plt.show() np.savez( os.path.join(SAVE_DIR, "electric_energy.npz"), t_sl=np.array(t_actual_energy), E_sl=np.array(E_sl_energy), t_pinn=t_arr, E_pinn=np.array(E_pinn_energy), gamma_th=gamma_th, E0_ref=E0_ref, ) # ── Préparation données pour plots 2D ──────────────────────────────────── n_visu = 128 x_visu = np.linspace(0.0, float(X_L), n_visu, endpoint=False) v_visu = np.linspace(-V_MAX, V_MAX, n_visu, endpoint=False) dv_visu = 2.0 * V_MAX / n_visu xx_pinn, vv_pinn = np.meshgrid(x_visu, v_visu) x_flat = jnp.array(xx_pinn.ravel()[:, None]) v_flat = jnp.array(vv_pinn.ravel()[:, None]) xv_flat = jnp.concatenate([x_flat, v_flat], axis=-1) M_v_pinn = np.exp(-0.5 * vv_pinn**2) / np.sqrt(2.0 * np.pi) xx_sl, vv_sl_mesh = np.meshgrid(x_sl, v_sl) M_v_sl = np.exp(-0.5 * vv_sl_mesh**2) / np.sqrt(2.0 * np.pi) _eval_f = jax.jit(jax.vmap(lambda xvt: sp.compute_density_inference(sp, xvt))) kx_sl_full = jnp.fft.fftfreq(NX_SL) * NX_SL * 2.0 * jnp.pi / float(X_L) kx_sl_s = jnp.where(kx_sl_full != 0, kx_sl_full, 1.0) nt = len(t_actual_diag) f_pinns, df_pinns = [], [] f_sls, df_sls = [], [] rho_pinns, E_pinns = [], [] rho_sls, E_sls = [], [] for idx, t_val in enumerate(t_actual_diag): t_col = jnp.full((n_visu * n_visu, 1), t_val) f_p = np.array(_eval_f(jnp.concatenate([xv_flat, t_col], axis=-1))).reshape( n_visu, n_visu ) f_pinns.append(f_p) df_pinns.append(f_p - M_v_pinn) rho_p = f_p.sum(axis=0) * dv_visu rho_pinns.append(rho_p) _, E_p = solve_poisson_fft(compute_rho_inference(sp, t_val)) E_pinns.append(np.array(E_p)) f_sl_i, rho_sl_i, _ = sl_diag[idx] f_sl_arr = np.array(f_sl_i).T f_sls.append(f_sl_arr) df_sls.append(f_sl_arr - M_v_sl) rho_sls.append(np.array(rho_sl_i)) rho_hat_sl = jnp.fft.fft(jnp.array(rho_sl_i) - 1.0) E_hat_sl = jnp.where(kx_sl_full != 0, -1j * rho_hat_sl / kx_sl_s, 0.0 + 0.0j) E_sls.append(np.array(jnp.fft.ifft(E_hat_sl).real)) # ── Sauvegarde des champs 2D / 1D pour post-traitement ──────────────────── np.savez( os.path.join(SAVE_DIR, "fields_diag.npz"), t_diag=np.array(t_actual_diag), # grilles PINN x_visu=x_visu, v_visu=v_visu, xx_pinn=xx_pinn, vv_pinn=vv_pinn, x_poisson=np.array(x_poisson), # grilles SL x_sl=x_sl, v_sl=v_sl, xx_sl=xx_sl, vv_sl=vv_sl_mesh, # champs PINN f_pinns=np.array(f_pinns), df_pinns=np.array(df_pinns), rho_pinns=np.array(rho_pinns), E_pinns=np.array(E_pinns), # champs SL f_sls=np.array(f_sls), df_sls=np.array(df_sls), rho_sls=np.array(rho_sls), E_sls=np.array(E_sls), ) print(f"Data saved in {SAVE_DIR}") # ── Plot 1 : f et δf ───────────────────────────────────────────────────── vmax_f = max(max(f.max() for f in f_pinns), max(f.max() for f in f_sls)) lev_f = np.linspace(0.0, vmax_f, 40) lim_df = max( max(np.abs(df).max() for df in df_pinns), max(np.abs(df).max() for df in df_sls), 1e-12, ) lev_df = np.linspace(-lim_df, lim_df, 41) fig1, axes1 = plt.subplots(4, nt, figsize=(4.5 * nt, 14), squeeze=False) row_defs = [ (f_pinns, "turbo", lev_f, None, xx_pinn, vv_pinn), (df_pinns, "RdBu_r", lev_df, "%.2e", xx_pinn, vv_pinn), (f_sls, "turbo", lev_f, None, xx_sl, vv_sl_mesh), (df_sls, "RdBu_r", lev_df, "%.2e", xx_sl, vv_sl_mesh), ] row_labels = [ r"PINN-FFT $f$", r"PINN-FFT $\delta f$", r"SL $f$", r"SL $\delta f$", ] for row, (data_list, cmap, lev, fmt, xx, vv) in enumerate(row_defs): for col in range(nt): ax = axes1[row, col] im = ax.contourf( xx, vv, data_list[col], levels=lev, cmap=cmap, extend="both" ) ax.contour( xx, vv, data_list[col], levels=lev[1:-1:4], colors="k", linewidths=0.4, alpha=0.5, ) cb = plt.colorbar(im, ax=ax) if fmt: cb.formatter.set_powerlimits((0, 0)) cb.update_ticks() if row == 3: ax.set_xlabel(r"$x$") if row == 0: ax.set_title(rf"$t = {t_actual_diag[col]:.2f}$") axes1[row, 0].set_ylabel(row_labels[row] + r"$\quad v$") plt.tight_layout() fig1.subplots_adjust(left=0.12) _save_fig(fig1, "f_and_df") plt.show() # ── Plot 2 : ρ−1 et E ──────────────────────────────────────────────────── fig2, axes2 = plt.subplots(2, nt, figsize=(4.5 * nt, 7), squeeze=False) for col, t_val in enumerate(t_actual_diag): ax = axes2[0, col] ax.plot(x_sl, rho_sls[col] - 1.0, "steelblue", lw=1.5, label="SL") ax.plot( x_visu, rho_pinns[col] - 1.0, "--", color="tomato", lw=1.5, label="PINN" ) ax.axhline(0.0, color="gray", ls=":", lw=0.8) ax.set_title(rf"$t = {t_val:.2f}$") if col == 0: ax.set_ylabel(r"$\rho - 1$") ax.legend() ax = axes2[1, col] ax.plot(x_sl, E_sls[col], "steelblue", lw=1.5, label="SL") ax.plot( np.array(x_poisson), E_pinns[col], "--", color="tomato", lw=1.5, label="PINN", ) ax.axhline(0.0, color="gray", ls=":", lw=0.8) ax.set_xlabel(r"$x$") if col == 0: ax.set_ylabel(r"$E$") ax.legend() plt.tight_layout() _save_fig(fig2, "rho_and_E") plt.show()