# %% """ Vlasov-Poisson 1D — flot PINN + Poisson FFT avec cache E mis à jour par epoch. Schéma Picard alterné : à chaque epoch : 1. calcul E_grids avec params courants → space._E_grids_cache (leaf pytree) 2. optimisation du flot avec E_grids figé Optimisations : - _E_grids_cache comme leaf du pytree JAX - compute_rho / compute_density_training : composition complète des flots 0..i (φ_i le plus interne, flots précédents gelés via stop_gradient) - warm_start_from_previous_flow - compute_all_energies : lax.map sur tous les instants, jit une seule fois """ import copy import os import sys import timeit import jax import jax.numpy as jnp import matplotlib.pyplot as plt import numpy as np from tqdm import tqdm 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 ( DEFAULT_N_BC_COLLOC, DEFAULT_N_IC_COLLOC, 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.networks.structure_preserving_nets.symplectic_nets import ( # noqa: E501 PeriodicGSymplecticNet, ) 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, ) from scimba_torch.utils.environment import ( get_static_terminal_width, is_static_width_environment, ) jax.config.update("jax_enable_x64", True) # ───────────────────────────────────────────────────────────────────────────── # Paramètres physiques # ───────────────────────────────────────────────────────────────────────────── CASE = "two_stream" 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 = 400 nb_models = 4 Tf = 24 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 = 40 # Réseau de flot : "mlp" (baseline) ou "sympnet" (symplectique, det J = 1 exact). NET = "sympnet" # "mlp" | "sympnet" MLP_HIDDEN = [30, 30] # archi MLP SYMP_WIDTH = 8 # largeur des couches du sympnet SYMP_NB_LAYERS = 13 # nombre de couches symplectiques composées SYMP_IDENTITY_INIT = True # init le sympnet ≈ identité (équiv. symplectique de add_id) SYMP_INIT_SCALE = 1e-2 # facteur sur scaling.weight : petit mais NON nul (ENG !) 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) # ───────────────────────────────────────────────────────────────────────────── # Style des figures (« LaTeX ») + dossier de sauvegarde paramétré # ───────────────────────────────────────────────────────────────────────────── # Police « LaTeX » : si une distribution LaTeX est installée, mettre VP_USE_TEX=1 # pour un vrai rendu usetex ; sinon 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", # cmr10 est une police MATH (pas de δ, φ, ρ, «—», «−» en texte brut → # « Glyph missing from font cmr10 »). STIXGeneral a l'allure Computer # Modern ET une couverture Unicode complète. mathtext reste en « cm ». "font.serif": ["STIXGeneral", "DejaVu Serif", "Computer Modern Roman"], "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 de sortie : le nom encode le cas test, la grille Poisson (espace + # temps), le nombre de flots, le nombre d'époques et le temps final, pour que # chaque configuration ait son propre dossier de figures. RUN_TAG = ( f"{CASE}_{NET}" f"_nx{NX_POISSON}_nq{NQ_POISSON}_nt{N_T_BINS}" f"_nm{nb_models}_ep{N_EPOCHS}_Tf{Tf:g}" ) SAVE_DIR = os.path.join( os.path.dirname(os.path.abspath(__file__)), "results_vp_fft", RUN_TAG ) os.makedirs(SAVE_DIR, exist_ok=True) print(f"Figures saved in {SAVE_DIR}") 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}")) # ───────────────────────────────────────────────────────────────────────────── # Poisson FFT # ───────────────────────────────────────────────────────────────────────────── 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) # ───────────────────────────────────────────────────────────────────────────── # 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 — utilise le cache figé + seulement φ_i # ───────────────────────────────────────────────────────────────────────────── 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 (indispensable au sympnet : identité # exacte à t_local=0). Flot courant : t_pt ; flots précédents : fin de # leur intervalle → t_local = Δt_k. 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 = solve_poisson_fft(compute_rho(space, t_scalar)) return E 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 # ───────────────────────────────────────────────────────────────────────────── class VPFlowSpace(DensityFlowFieldsApproximationSpace): """ Étend DensityFlowFieldsApproximationSpace avec : 1. _E_grids_cache : leaf pytree → mises à jour vues par JAX 2. compute_density_training : compose les flots 0..i (φ_i le plus interne) """ def __init__(self, *args, n_t_bins: int = N_T_BINS, **kwargs): super().__init__(*args, **kwargs) self._E_grids_cache = jnp.zeros((n_t_bins, NX_POISSON)) self._n_t_bins = n_t_bins # ── Pytree ─────────────────────────────────────────────────────────────── # Plus de tree_flatten/tree_unflatten ici : la classification generique de # ScimbaPytree range deja `_E_grids_cache` (un jnp.ndarray) dans les # children et `_n_t_bins` (un int) dans l'aux_data. Le cache n'est pas # optimise parce qu'il n'est PAS declare `trainable`, pas parce qu'il # serait range ailleurs -- c'est `auto_partition` qui tranche. # ── 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) # ── Déterminant du jacobien ─────────────────────────────────────────────── def _abs_det_jac_wrt_x(self, model, xv, extra): """|det ∂φ/∂(x,v)|. Pour un réseau symplectique (PeriodicGSymplecticNet), le déterminant vaut 1 EXACTEMENT : on utilise son `abs_det_jacobian` analytique (comme DensityFlowInvertibleApproximationSpace dans lagrangian_pinn_linear_vlasov). Ça évite le jacrev numérique du parent (coûteux, et qui dérive à travers le `q % period` du sympnet). Le MLP n'a pas cette méthode → on retombe sur le jacrev numérique du parent. """ if hasattr(model, "abs_det_jacobian"): return model.abs_det_jacobian(jnp.concatenate([xv, extra], axis=-1)) return super()._abs_det_jac_wrt_x(model, xv, extra) # ── compute_density en mode entraînement ───────────────────────────────── def compute_density_training(self, space, *args): """ Mode entraînement : compose les flots 0..i sans lax.switch. Utilisé dans batched_fn (points de collocation). """ 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] # temps local t - t_start du flot k (cf. sympnet : identité à t_local=0) t_local = t_k - space.time_intervals[k, 0:1] extra = jnp.concatenate([t_local, 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 # ── Cache E ─────────────────────────────────────────────────────────────── def update_e_cache(self, t_bins: jnp.ndarray): self.inference_mode = False E_grids = jax.lax.map(lambda t: compute_e_grid(self, t), t_bins) self._E_grids_cache = jax.lax.stop_gradient(E_grids) # `flow_partition`/`flow_combine` ont disparu : elles gelaient UNE feuille # (`_E_grids_cache`) et laissaient tout le reste actif. Le defaut du projecteur # est maintenant `auto_partition`, qui fait l'inverse -- tout est gele sauf ce # qui est declare `trainable`, c'est-a-dire le flot courant. Le cache reste # hors de la Gram, et les flots deja entraines aussi. # # ───────────────────────────────────────────────────────────────────────────── # # Tests unitaires # # ───────────────────────────────────────────────────────────────────────────── # def test_poisson_fft(): # print("\n" + "=" * 60) # print("=== Étape 1 : solve_poisson_fft ===") # 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 ✗'}" # ) # def loss_e(rho): # _, E = solve_poisson_fft(rho) # return jnp.sum(E**2) # grad_rho = jax.grad(loss_e)(rho_test) # has_nan = bool(jnp.any(jnp.isnan(grad_rho))) # max_grad = float(jnp.abs(grad_rho).max()) # print( # f" max|∂loss/∂rho| = {max_grad:.2e} " # f"{'NaN ✗' if has_nan else ('OK ✓' if max_grad > 1e-10 else 'COUPÉ ✗')}" # ) # 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_rho - grad_fd).max()) # print( # f" max|grad_AD-grad_FD| = {err_fd:.2e} {'OK ✓' if err_fd < 1e-5 else 'FAIL ✗'}" # ) # def test_catmull_rom(): # print("\n" + "=" * 60) # print("=== Étape 2 : interpolation Catmull-Rom ===") # 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 loss_c(field): # return jnp.sum(catmull_rom_interp(field, x_fine) ** 2) # grad_field = jax.grad(loss_c)(E_grid) # has_nan = bool(jnp.any(jnp.isnan(grad_field))) # max_grad = float(jnp.abs(grad_field).max()) # print( # f" max|∂loss/∂field| : {max_grad:.2e} " # f"{'NaN ✗' if has_nan else ('OK ✓' if max_grad > 1e-10 else 'COUPÉ ✗')}" # ) # def test_compute_rho(space): # print("\n" + "=" * 60) # print("=== Étape 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"(réseau non entraîné {'OK ✓' if err < 0.1 else 'FAIL ✗'})" # ) # mass = float(rho_num.sum() * dx_poisson) # print( # f" ∫ρ dx = {mass:.6f} (attendu {float(X_L):.6f}) " # f"{'OK ✓' if abs(mass - float(X_L)) < 0.1 else 'FAIL ✗'}" # ) # def loss_rho(s): # return jnp.sum(compute_rho(s, 0.0) ** 2) # grad = jax.grad(loss_rho)(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 params|| = {max_grad:.2e} " # f"{'NaN ✗' if has_nan else ('OK ✓' if max_grad > 1e-10 else 'COUPÉ ✗')}" # ) # def test_compute_e(space): # print("\n" + "=" * 60) # print("=== Étape 4 : compute_E ===") # space.idx_current_flow = 0 # space.inference_mode = False # E_num = compute_e_grid(space, 0.0) # E_exact = (EPS / K) * jnp.sin(K * x_poisson) # has_nan = bool(jnp.any(jnp.isnan(E_num))) # corr = float( # jnp.dot(E_num, E_exact) # / (jnp.linalg.norm(E_num) * jnp.linalg.norm(E_exact) + 1e-16) # ) # print( # f" corrélation E(t=0) / E_exact = {corr:.4f} " # f"{'NaN ✗' if has_nan else '(réseau non entraîné)'}" # ) # print( # f" max|E_num| = {float(jnp.abs(E_num).max()):.4e} " # f"(attendu {float(EPS / K):.4e})" # ) # def loss_e(s): # return jnp.sum(compute_e_grid(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 params|| = {max_grad:.2e} " # f"{'NaN ✗' if has_nan else ('OK ✓' if max_grad > 1e-10 else 'COUPÉ ✗')}" # ) # def test_e_cache(space): # print("\n" + "=" * 60) # print("=== Étape 5 : E_grids_cache comme leaf pytree ===") # space.idx_current_flow = 0 # space.inference_mode = False # t_bins = jnp.linspace(*time_intervals[0], N_T_BINS) # space.update_e_cache(t_bins) # has_nan = bool(jnp.any(jnp.isnan(space._E_grids_cache))) # print( # f" E_grids_cache shape : {space._E_grids_cache.shape} " # f"(attendu ({N_T_BINS}, {NX_POISSON}))" # ) # print(f" NaN dans cache : {'OUI ✗' if has_nan else 'non ✓'}") # print(f" max|E_cache| : {float(jnp.abs(space._E_grids_cache).max()):.4e}") # leaves = jax.tree_util.tree_leaves(space) # cache_found = any( # hasattr(ell, "shape") and ell.shape == (N_T_BINS, NX_POISSON) for ell in leaves # ) # print(f" Cache dans pytree : {'OUI ✓' if cache_found else 'NON ✗'}") # 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. # Avant correctif, compute_rho appliquait φ_i en dernier → écart non nul ici. # """ # print("\n" + "=" * 60) # print("=== Étape 6 : ordre de composition compute_rho (idx=1) ===") # if len(time_intervals) < 2: # print(" (un seul flot — test sauté)") # 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 - réf. pleine compo| = {err:.2e} " # f"{'OK ✓' if err < 1e-9 else 'FAIL ✗'}" # ) # def loss_rho(s): # return jnp.sum(compute_rho(s, t_test) ** 2) # grad = jax.grad(loss_rho)(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 params|| : {max_grad:.2e} " # f"{'NaN ✗' if has_nan else ('OK ✓' if max_grad > 1e-10 else 'COUPÉ ✗')}" # ) # space.idx_current_flow = 0 # # ───────────────────────────────────────────────────────────────────────────── # Résidu Vlasov # ───────────────────────────────────────────────────────────────────────────── class BatchedParamVecFunction(ParamVecFunction): def vmap_on_physical_variables(self): return self def compose_post_processing_with(self, post_processing): result = super().compose_post_processing_with(post_processing) return BatchedParamVecFunction( size=result.size, dim=result.dims, fn=result.fn, f_type=result.f_type, ) class VlasovResidualFFT(InteriorResidual): def __init__(self, domain, time_domain: tuple[float, float], f_rhs=None): super().__init__( domain=domain, size=2, model_type="x_v_t", f_rhs=f_rhs, time_domain=time_domain, ) 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 = phi_x.dims f_type = phi_x.f_type dt_phi_x_b = phi_x.d_t().vmap_on_physical_variables() dt_phi_v_b = phi_v.d_t().vmap_on_physical_variables() dx_phi_x_b = phi_x.gradient("x").vmap_on_physical_variables() dx_phi_v_b = phi_v.gradient("x").vmap_on_physical_variables() dv_phi_x_b = phi_x.gradient("v").vmap_on_physical_variables() dv_phi_v_b = phi_v.gradient("v").vmap_on_physical_variables() t0_ = self._t0 t1_ = self._t1 def batched_fn(space, x_batch, v_batch, t_batch): x_2d = x_batch.reshape(-1, 1) v_2d = v_batch.reshape(-1, 1) t_2d = t_batch.reshape(-1, 1) t_1d = t_2d[:, 0] # E figé pour ce pas d'optim (vrai Picard) : stop_gradient AU POINT # DE LECTURE, sinon l'optimiseur traite _E_grids_cache (leaf pytree) # comme un paramètre entraînable et « triche » en bougeant E au lieu # du flot. Le stop_gradient au stockage (update_e_cache) ne suffit pas. E_grids = jax.lax.stop_gradient(space._E_grids_cache) n_t_bins = E_grids.shape[0] dt_bin = (t1_ - t0_) / (n_t_bins - 1) idx_f = (t_1d - t0_) / dt_bin idx0 = jnp.clip(jnp.floor(idx_f).astype(jnp.int32), 0, n_t_bins - 2) alpha = idx_f - idx0 E_grid_i = (1.0 - alpha[:, None]) * E_grids[idx0] + alpha[ :, None ] * E_grids[idx0 + 1] x_1d = x_2d[:, 0] E_vals = jax.vmap(lambda eg, xq: catmull_rom_interp(eg, xq.reshape(1))[0])( E_grid_i, x_1d ) dt_px = dt_phi_x_b(space, x_2d, v_2d, t_2d).ravel() dt_pv = dt_phi_v_b(space, x_2d, v_2d, t_2d).ravel() dx_px = dx_phi_x_b(space, x_2d, v_2d, t_2d).ravel() dx_pv = dx_phi_v_b(space, x_2d, v_2d, t_2d).ravel() dv_px = dv_phi_x_b(space, x_2d, v_2d, t_2d).ravel() dv_pv = dv_phi_v_b(space, x_2d, v_2d, t_2d).ravel() v_vals = v_2d[:, 0] res_x = dt_px + v_vals * dx_px + E_vals * dv_px res_v = dt_pv + v_vals * dx_pv + E_vals * dv_pv return jnp.stack([res_x, res_v], axis=-1) return BatchedParamVecFunction( size=2, dim=dims, fn=batched_fn, f_type=f_type, ) # ───────────────────────────────────────────────────────────────────────────── # CI du flot # ───────────────────────────────────────────────────────────────────────────── 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, ) def construct_residual(self, *args) -> PARAM_FUNC_TYPE: return args[0].set_t_0(self.time_domain[0]) # ───────────────────────────────────────────────────────────────────────────── # 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) # ───────────────────────────────────────────────────────────────────────────── # Projector avec mise à jour du cache avant chaque epoch # ───────────────────────────────────────────────────────────────────────────── class ProjectorWithECache(Projector): def __init__(self, *args, update_e_cache_fn=None, **kwargs): super().__init__(*args, **kwargs) self._update_e_cache_fn = update_e_cache_fn def project( self, key, space, n_epochs, n_colloc, n_bc_colloc=DEFAULT_N_BC_COLLOC, n_ic_colloc=DEFAULT_N_IC_COLLOC, verbose=False, **kwargs, ): init_best_loss = self.losses.get_infinity_losses() losses = self.losses.set_initial_losses_history(n_epochs) one_step_optim = self.build_one_step_optim(n_colloc, n_bc_colloc, n_ic_colloc) new_space = space best_space = space best_loss = init_best_loss optimizer = self.optimizer tqdm_ncols = get_static_terminal_width() tqdm_dynamic = not is_static_width_environment() tqdm_disable = verbose or (tqdm_ncols == 0) tqdm_position = kwargs.get("tqdm_position", 0) tqdm_desc = kwargs.get("tqdm_desc", "Training") tqdm_leave = kwargs.get("tqdm_leave", "True") loop = tqdm( total=n_epochs, desc=tqdm_desc, bar_format="{l_bar}{bar}| {n_fmt}/{total_fmt}" "[{elapsed}<{remaining}] {postfix}", disable=tqdm_disable, leave=tqdm_leave, position=tqdm_position, ascii=" |", file=sys.stdout, dynamic_ncols=tqdm_dynamic, ncols=tqdm_ncols, ) for i in range(n_epochs): if self._update_e_cache_fn is not None: new_space = self._update_e_cache_fn(new_space) new_loss, new_space, key, optimizer = one_step_optim( new_space, key, optimizer ) if new_loss["total"] < best_loss["total"]: best_loss, best_space = new_loss, new_space if verbose: print("Epoch %d: New best loss: %.2e" % (i, new_loss["total"])) losses = losses.update_losses_history(i, new_loss) # update loop informations # rate_epoch = int(math.floor(max_rate * epoch/self.max_epochs)) postfix_str = "loss: %.1e -> %.1e" % ( losses.losses_history["total"][1], best_loss["total"], ) loop.set_postfix_str(postfix_str) loop.update(1) loop.refresh() loop.close() children, aux_data = self.tree_flatten() # ⚠ Recollage POSITIONNEL des children du Projector. Le commentaire qui # etait ici annoncait 5 children ; il y en a 6 depuis que `key` a ete # ajoute, et rien n'a leve : `tree_unflatten` lisait `children[5]` hors # tuple. On reconstruit donc a partir de `children` au lieu de reecrire # la liste, pour qu'un champ de plus ne casse plus rien en silence. nchildren = ( children[0], best_space, losses, optimizer, best_loss, *children[5:], ) new_projector = Projector.tree_unflatten(aux_data, nchildren) return key, new_projector # init_best_loss = self.losses.get_infinity_losses() # losses_history = self.losses.get_initial_losses_history(n_epochs) # one_step_optim = self.build_one_step_optim(n_colloc, n_bc_colloc, n_ic_colloc) # new_space = space # best_space = space # best_loss = init_best_loss # optimizer = self.optimizer # tqdm_ncols = get_static_terminal_width() # tqdm_dynamic = not is_static_width_environment() # tqdm_disable = verbose or (tqdm_ncols == 0) # loop = tqdm( # total=n_epochs, # desc=kwargs.get("tqdm_desc", "Training"), # bar_format="{l_bar}{bar}| {n_fmt}/{total_fmt}[{elapsed}<{remaining}] {postfix}", # disable=tqdm_disable, # leave=kwargs.get("tqdm_leave", True), # position=kwargs.get("tqdm_position", 0), # ascii=" |", # file=sys.stdout, # dynamic_ncols=tqdm_dynamic, # ncols=tqdm_ncols, # ) # for i in range(n_epochs): # if self._update_e_cache_fn is not None: # new_space = self._update_e_cache_fn(new_space) # new_loss, new_space, key, optimizer = one_step_optim( # new_space, key, optimizer # ) # if new_loss["total"] < best_loss["total"]: # best_loss, best_space = new_loss, new_space # losses_history = self.losses.update_losses_history(i, new_loss) # loop.set_postfix_str( # "loss: %.1e -> %.1e" % (losses_history["total"][1], best_loss["total"]) # ) # loop.update(1) # loop.refresh() # loop.close() # self.optimizer = optimizer # children, aux_data = self.tree_flatten() # nchildren = (children[0], best_space, self.losses, optimizer) # aux_data["best_loss"] = best_loss # new_projector = Projector.tree_unflatten(aux_data, nchildren) # return key, new_projector # ───────────────────────────────────────────────────────────────────────────── # 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) # Le sympnet voit (x, v) comme (q, p) et t comme entrée conditionnelle ; sa # périodicité en q=x est gérée par period=(X_L,). add_id DOIT être False pour le # sympnet (sinon nn(.) + (x,v) n'est plus symplectique). if NET == "mlp": flow_models = [ MLP( in_size=dim_x + dim_v + 1, hidden_sizes=MLP_HIDDEN, out_size=dim_x + dim_v, activation="tanh", key=keys[i], embedding="periodic", periods=(float(X_L),), embedding_axes=(0,), ) for i in range(nb_models) ] elif NET == "sympnet": def init_sympnet_to_identity(net, scale=SYMP_INIT_SCALE): """Rapproche le sympnet de l'identité en réduisant tous les `scaling.weight`. Chaque couche est un cisaillement p += h·(scaling(mu)·act(...))·W ; le terme est *gated* par scaling(mu). Multiplier scaling.weight par un petit `scale` rend le flot ≈ identité (perturbation O(scale)) — équivalent symplectique du add_id du MLP. IMPORTANT : on ne met PAS scaling à zéro. Avec ENG (gradient naturel), un zéro exact annule les gradients des autres poids (linear_q, linear_mu, …) -> matrice de Gram singulière -> pas géant -> explosion. Un `scale` petit mais non nul garde la Gram de plein rang. """ return jax.tree_util.tree_map_with_path( lambda path, leaf: ( leaf * scale if any(getattr(k, "name", None) == "scaling" for k in path) else leaf ), net, ) flow_models = [ PeriodicGSymplecticNet( size=dim_x + dim_v, conditional_size=1, width=SYMP_WIDTH, nb_layers=SYMP_NB_LAYERS, activation="tanh", period=(float(X_L),), key=keys[i], ) for i in range(nb_models) ] if SYMP_IDENTITY_INIT: flow_models = [init_sympnet_to_identity(m) for m in flow_models] else: raise ValueError(f"NET inconnu: {NET!r}") add_id = NET == "mlp" space = VPFlowSpace( dim=dim_x, velocity_dim=dim_v, params_dim=params_dim, models=flow_models, list_models_fields=[], add_id=add_id, 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 ÉTAPES 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_e_cache(space) # test_compose_order(space) # print("\n=== Tous les tests OK — passage à l'entraînement ===\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 per_flow_times = [] 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} ===") # Warm start depuis le flot précédent # if i > 0: # space.warm_start_from_previous_flow() sampler.renew_sampler(UniformTimeSampler(domain_t), "t") model.renew_time_domain(domain_t) space.inference_mode = False t_bins_i = jnp.linspace(float(domain_t[0]), float(domain_t[1]), N_T_BINS) # Initialise le cache E space.update_e_cache(t_bins_i) def make_update_fn(t_bins, flow_idx): def update_fn(sp): sp.idx_current_flow = flow_idx sp.inference_mode = False sp.update_e_cache(t_bins) return sp return update_fn update_fn = make_update_fn(t_bins_i, i) pinn = ProjectorWithECache( model, space, sampler, optimizer="ENG", weights=weights, one_loss_per_residual=True, matrix_regularization=1.0e-6, update_e_cache_fn=update_fn, ) if i == 0: print("Optimiseur : ENG") print(f"Network : {NET}") 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 doit composer avec les # flots précédents ENTRAÎNÉS, pas avec le `space` initial non entraîné. space = pinn.space list_of_pinns.append(copy.deepcopy(pinn)) pinn_total_time += elapsed per_flow_times.append(elapsed) print(f" time : {elapsed:.1f}s") print(f"\n=== PINN FFT: total time: {pinn_total_time:.1f}s ===") # ───────────────────────────────────────────────────────────────────────────── # params.txt — récapitulatif des paramètres du test (écrit dans SAVE_DIR) # ───────────────────────────────────────────────────────────────────────────── def _fmt_hms(s): """Secondes -> '~MmSSs'.""" m, sec = divmod(int(round(s)), 60) return f"~{m}m{sec:02d}s" def write_params_txt(): n_params = int(sum(jnp.size(x) for x in jax.tree_util.tree_leaves(flow_models[0]))) dt_slice = Tf / nb_models # largeur d'une tranche temporelle if NET == "mlp": arch_line = f"hidden sizes : {MLP_HIDDEN}" else: arch_line = ( f"sympnet width : {SYMP_WIDTH}\nsympnet nb_layers: {SYMP_NB_LAYERS}" ) device = str(jax.devices()[0]) lines = [ "=== Test parameters (Vlasov-Poisson FFT, flow-map PINN) ===", f"case : {CASE}", f"network : {NET}", "optimizer : ENG", "", f"nb_models : {nb_models}", f"slice width dt : {dt_slice:g} (Tf / nb_models)", arch_line, f"dofs per model : {n_params}", f"NQ_POISSON (v quadrature) : {NQ_POISSON}", f"NX_POISSON (Poisson grid) : {NX_POISSON}", f"N_T_BINS (E-cache, per slice): {N_T_BINS}", f"n_colloc : {N_COLLOC}", f"n_ic_colloc : {N_IC_COLLOC}", f"n_epochs : {N_EPOCHS}", f"Tf : {Tf:g}", f"V_MAX : {float(V_MAX):.2f}", "", f"device : {device}", f"total train time : {pinn_total_time:.0f} s ({_fmt_hms(pinn_total_time)}) " "(incl. JIT compile, all flows)", ] if per_flow_times: mean_t = pinn_total_time / len(per_flow_times) first, last = per_flow_times[0], per_flow_times[-1] lines.append( f"time / network : {mean_t:.0f} s ({_fmt_hms(mean_t)}) mean " f"[{first:.0f} s ({_fmt_hms(first)}) first .. " f"{last:.0f} s ({_fmt_hms(last)}) last; grows with composition depth]" ) txt = "\n".join(lines) + "\n" with open(os.path.join(SAVE_DIR, "params.txt"), "w") as fh: fh.write(txt) print(f"\n{txt}-> {os.path.join(SAVE_DIR, 'params.txt')}") write_params_txt() # ───────────────────────────────────────────────────────────────────────────── # 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 True: 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 ────────────────────────────────────────────────────────── 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) ax.set_title(f"Loss — flow {i + 1}, t ∈ {time_intervals[i]}") ax.set_xlabel("epoch") ax.legend() ax.grid(True, alpha=0.3) plt.tight_layout() _save_fig(fig, f"loss_flow_{i + 1}") plt.show() # ── 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 — lax.map sur tous les instants ─────────────────── 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_energy = plt.figure(figsize=(9, 4)) plt.semilogy(np.array(t_actual_energy), E_sl_energy, "steelblue", lw=2, label="SL") plt.semilogy(t_arr, E_pinn_energy, "tomato", lw=1.5, label="PINN FFT") plt.semilogy( t_arr, E0_ref * np.exp(2 * gamma_th * t_arr), "k--", lw=1, label=f"theory γ={gamma_th}", ) plt.xlabel("t") plt.ylabel(r"$\frac{1}{2}\int E^2\,dx$") plt.title(f"Landau damping NX={NX_POISSON} NQ={NQ_POISSON} NT={N_T_BINS}") plt.legend() plt.tight_layout() _save_fig(fig_energy, "electric_energy") plt.show() # ── 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)) # ── 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) fig1.suptitle("f and δf = f−M(v) — PINN-FFT vs SL", fontsize=13) 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 = ["PINN-FFT f", "PINN-FFT δf", "SL f", "SL δ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() ax.set_xlabel("x") ax.set_ylabel("v") if row == 0: ax.set_title(f"t = {t_actual_diag[col]:.3f}") axes1[row, 0].set_ylabel(row_labels[row]) for row, (label, color) in enumerate( [ ("PINN-FFT", "tomato"), ("PINN-FFT", "tomato"), ("SL", "steelblue"), ("SL", "steelblue"), ] ): axes1[row, 0].text( -0.18, 0.5, label, transform=axes1[row, 0].transAxes, va="center", ha="center", rotation=90, fontsize=11, fontweight="bold", color="white", bbox=dict( facecolor=color, edgecolor="none", boxstyle="round,pad=0.35", alpha=0.9 ), ) plt.tight_layout() fig1.subplots_adjust(left=0.12) _save_fig(fig1, "phase_space_f_df") plt.show() # ── Plot 2 : ρ−1 et E ──────────────────────────────────────────────────── fig2, axes2 = plt.subplots(2, nt, figsize=(4.5 * nt, 7), squeeze=False) fig2.suptitle("ρ−1 and E — PINN-FFT vs SL", fontsize=13) 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(f"ρ−1 t={t_val:.2f}") ax.set_xlabel("x") if col == 0: 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_title(f"E t={t_val:.2f}") ax.set_xlabel("x") if col == 0: ax.legend() plt.tight_layout() _save_fig(fig2, "rho_E") plt.show() # ── Coupes ─────────────────────────────────────────────────────────────── def _cut_v(x_val, label, name): j_p = int(np.argmin(np.abs(x_visu - x_val))) j_s = int(np.argmin(np.abs(x_sl - x_val))) fig, axes = plt.subplots(1, nt, figsize=(4.0 * nt, 4), squeeze=False) fig.suptitle(f"Cut {label} — PINN-FFT vs SL", fontsize=13) for col in range(nt): ax = axes[0, col] ax.plot( v_visu, f_pinns[col][:, j_p], color="tomato", lw=1.8, label="PINN-FFT" ) ax.plot( v_sl, f_sls[col][:, j_s], color="steelblue", lw=1.5, ls="--", label="SL" ) ax.set_title(f"t = {t_actual_diag[col]:.3f}") ax.set_xlabel("v") ax.grid(True, alpha=0.3) if col == 0: ax.set_ylabel(f"f({label}, v, t)") ax.legend() plt.tight_layout() _save_fig(fig, name) plt.show() def _cut_x(v_val, label, name): i_p = int(np.argmin(np.abs(v_visu - v_val))) i_s = int(np.argmin(np.abs(v_sl - v_val))) fig, axes = plt.subplots(1, nt, figsize=(4.0 * nt, 4), squeeze=False) fig.suptitle(f"Cut {label} — PINN-FFT vs SL", fontsize=13) for col in range(nt): ax = axes[0, col] ax.plot( x_visu, f_pinns[col][i_p, :], color="tomato", lw=1.8, label="PINN-FFT" ) ax.plot( x_sl, f_sls[col][i_s, :], color="steelblue", lw=1.5, ls="--", label="SL" ) ax.set_title(f"t = {t_actual_diag[col]:.3f}") ax.set_xlabel("x") ax.grid(True, alpha=0.3) if col == 0: ax.set_ylabel(f"f(x, {label}, t)") ax.legend() plt.tight_layout() _save_fig(fig, name) plt.show() _cut_v(np.pi * 1.0, "x=π", "cut_v_x_pi") _cut_v(np.pi * 2.0, "x=2π", "cut_v_x_2pi") _cut_v(np.pi * 3.0, "x=3π", "cut_v_x_3pi") _cut_x(0.0, "v=0", "cut_x_v0") # ── Flot arrière composé ───────────────────────────────────────────────── def plot_flow_components(t_vals=None, n_visu_flow=64): if t_vals is None: t_vals = t_actual_diag _eval_flow = jax.jit( jax.vmap( lambda xvt: sp.compute_flow_backward_compose_continuous_inference( sp, xvt ) ) ) x_f = np.linspace(0.0, float(X_L), n_visu_flow, endpoint=False) v_f = np.linspace(-V_MAX, V_MAX, n_visu_flow, endpoint=False) xx_f, vv_f = np.meshgrid(x_f, v_f) xv_f = jnp.concatenate( [ jnp.array(xx_f.ravel()[:, None]), jnp.array(vv_f.ravel()[:, None]), ], axis=-1, ) nt_f = len(t_vals) fig, axes = plt.subplots(2, nt_f, figsize=(4.0 * nt_f, 7), squeeze=False) fig.suptitle( "Composed backward flow — δφ_x = φ_x−x | δφ_v = φ_v−v", fontsize=13 ) for col, t_val in enumerate(t_vals): t_col = jnp.full((n_visu_flow * n_visu_flow, 1), t_val) flow = np.array(_eval_flow(jnp.concatenate([xv_f, t_col], axis=-1))) dphi_x = flow[:, 0].reshape(n_visu_flow, n_visu_flow) - xx_f dphi_v = flow[:, 1].reshape(n_visu_flow, n_visu_flow) - vv_f for row, (data, row_label) in enumerate( [(dphi_x, "δφ_x"), (dphi_v, "δφ_v")] ): ax = axes[row, col] vmax = max(float(np.abs(data).max()), 1e-10) im = ax.contourf( xx_f, vv_f, data, levels=40, cmap="RdBu_r", vmin=-vmax, vmax=vmax ) ax.contour( xx_f, vv_f, data, levels=10, colors="k", linewidths=0.3, alpha=0.4 ) plt.colorbar(im, ax=ax, format="%.2e") ax.set_xlabel("x") if col == 0: ax.set_ylabel(f"v [{row_label}]") if row == 0: ax.set_title(f"t = {t_val:.2f}") plt.tight_layout() _save_fig(fig, "flow_components") plt.show() plot_flow_components()