# %% """Vlasov-Poisson 1D nonlinéaire — amortissement de Landau. ∂_t f + v ∂_x f + E ∂_v f = 0 ∂_x E(x,t) = ρ(x,t) - 1, ρ = ∫ f dv f(x,v,0) = f_0(x,v) = (1 + ε cos(kx)) M(v) Approche Lagrangienne : on cherche le flot arrière φ = (φ_x, φ_v) tel que f(x,v,t) = f_0(φ(x,v,t)) |det J_φ| et φ vérifie l'équation caractéristique : ∂_t φ + v ∂_x φ + E(x,t) ∂_v φ = 0, φ(x,v,0) = (x,v) """ # %% import copy import timeit import jax import jax.numpy as jnp import matplotlib.pyplot as plt from scimba_jax.domains.meshless_domains.base import VolumetricDomain 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 ( NDARRAYS_FUNC_TYPE, PARAM_FUNC_TYPE, InitialResidual, InteriorResidual, ) jax.config.update("jax_enable_x64", True) # ───────────────────────────────────────────────────────────────────────────── # Paramètres physiques (amortissement de Landau linéaire) # ───────────────────────────────────────────────────────────────────────────── CASE = "landau_damping" # "landau_damping" ou "two_stream" K = 0.5 # nombre d'onde X_L = 2.0 * jnp.pi / K # longueur du domaine périodique ≈ 4π if CASE == "landau_damping": EPS = 0.1 # amplitude de la perturbation V_MAX = 6.0 # borne en vitesse elif CASE == "two_stream": EPS = 0.05 V_MAX = 5 * jnp.pi V0 = 3.0 else: raise ValueError(f"Case {CASE!r} not implemented") def f0(xv, *_, case="landau_damping"): """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) # ───────────────────────────────────────────────────────────────────────────── # Solveur Semi-Lagrangien Vlasov-Poisson 1D (référence numérique) # Splitting de Strang : x(dt/2) → Poisson → v(dt) → x(dt/2) # Transport exact par décalage spectral (FFT), périodique en x et v. # ───────────────────────────────────────────────────────────────────────────── NX_SL = 512 NV_SL = 512 DT_SL = 0.002 def run_sl_vlasov_poisson(t_checkpoints, Nx=NX_SL, Nv=NV_SL, dt=DT_SL): """SL Vlasov-Poisson périodique en x — résultats à plusieurs instants. Returns : (results, t_actual, x_grid, v_grid) results : list of (f, rho, phi) à chaque checkpoint (f.shape=(Nx,Nv)) t_actual : temps effectifs alignés sur les pas dt """ 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") # (Nx, Nv) xv = jnp.stack([xx, vv], axis=-1) f = f0(xv, case=CASE) # Vecteurs d'onde kx = jnp.fft.fftfreq(Nx) * Nx * 2.0 * jnp.pi / float(X_L) # (Nx,) kv = jnp.fft.fftfreq(Nv) * Nv * 2.0 * jnp.pi / (2.0 * V_MAX) # (Nv,) def compute_phi_e(f): rho = jnp.sum(f, axis=1) * dv_sl # (Nx,) rho_hat = jnp.fft.fft(rho - 1.0) phi_hat = jnp.where(kx != 0, rho_hat / kx**2, 0.0 + 0.0j) E_hat = -1j * kx * phi_hat return rho, jnp.fft.ifft(phi_hat).real, jnp.fft.ifft(E_hat).real def advect_x(f, dt_sub): f_hat = jnp.fft.fft(f, axis=0) # (Nx, Nv) phase = jnp.exp(-1j * kx[:, None] * v_grid[None, :] * dt_sub) return jnp.fft.ifft(f_hat * phase, axis=0).real def advect_v(f, E, dt_sub): f_hat = jnp.fft.fft(f, axis=1) # (Nx, Nv) phase = jnp.exp(-1j * kv[None, :] * E[:, None] * dt_sub) return jnp.fft.ifft(f_hat * phase, 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) # Poisson sur l'état final (output + prochain E) return f, rho, phi # Aligner chaque checkpoint sur le pas dt le plus proche (arrondi standard) 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 : {n_steps} pas (dt={dt}, t_end={n_steps * dt:.3f})") 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) rho = phi = jnp.zeros(Nx) for i in range(n_steps): f, rho, phi = step(f) step_num = i + 1 if step_num in step_to_idxs: for idx in step_to_idxs[step_num]: 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 # ───────────────────────────────────────────────────────────────────────────── # Résidus physiques # ───────────────────────────────────────────────────────────────────────────── class VlasovPoissonResidual(InteriorResidual): """Résidu combiné Vlasov + Poisson (taille 3). Variables reçues dans l'ordre de create_variables : inv_phi — flot arrière T^{-1} : (x,v,t) → (x_0,v_0) [taille 2] densite_f — densité cinétique f(x,v,t) [scalaire] phi_e — potentiel électrique φ_E(x,t) [scalaire] rho — moment 0 : ρ(x,t) = ∫ f dv [scalaire] u — moment 1 : u(x,t) = ∫ v f dv / ρ [vecteur] T — moment 2 : température T(x,t) [scalaire] Équations : Vlasov : ∂_t φ_i + v ∂_x φ_i − ∂_x φ_E ∂_v φ_i = 0 (i = x, v) Poisson : −∂_xx φ_E = (ρ − 1)/ε (φ_E = φ/ε, φ potentiel physique) """ def __init__(self, domain: VolumetricDomain, time_domain, f_rhs=None): super().__init__( domain=domain, size=3, # 2 composantes Vlasov + 1 Poisson model_type="x_v_t", f_rhs=f_rhs, time_domain=time_domain, ) def construct_residual( self, inv_phi: PARAM_FUNC_TYPE, # flot arrière, taille 2 _densite_f: PARAM_FUNC_TYPE, # densité cinétique f(x,v,t) [non utilisé ici] phi_e: PARAM_FUNC_TYPE, # potentiel électrique φ_E(x,t) rho: PARAM_FUNC_TYPE, # moment 0 ρ(x,t) _u: PARAM_FUNC_TYPE, # moment 1 u(x,t) [non utilisé ici] _T: PARAM_FUNC_TYPE, # moment 2 T(x,t) [non utilisé ici] ) -> PARAM_FUNC_TYPE: phi_x, phi_v = inv_phi.components() # ── dérivées temporelles ───────────────────────────────────────────── dt_phi_x = phi_x.d_t() dt_phi_v = phi_v.d_t() # ── E = −∂_x φ_E (force sur les caractéristiques) ────────────────── dphi_e_dx = phi_e.gradient("x") # ParamFieldFunction, taille 1 # ── termes d'advection : v ∂_x φ_i + E ∂_v φ_i ────────────────────── def _advection(phi_comp): dphi_dx = phi_comp.gradient("x") # ∂_x φ_i — ParamFieldFunction dphi_dv = phi_comp.gradient("v") # ∂_v φ_i — ParamFieldFunction # réseau apprend φ̃ = φ/ε → E = −∂_x φ_E = −ε ∂_x φ̃ def v_func(v): return v # fonction identité pour v E_val = -EPS * dphi_e_dx adv_x = dphi_dx * v_func adv_v = dphi_dv * E_val return adv_x + adv_v adv_x = _advection(phi_x) adv_v = _advection(phi_v) # ── résidu de Poisson : −∂_xx φ_E − (ρ − 1) = 0 ──────────────────── # phi_e est un ParamScalarFunction avec f_type "x_t" # laplacian("x") = ∂_xx φ_E (scalaire, 1D donc = ∂²/∂x²) lap_phi_e = phi_e.laplacian("x") # ParamScalarFunction poisson_res = lap_phi_e + (rho - 1.0) / EPS # résidu : ∂_xx φ_E + (ρ-1)/ε = 0 return ParamVecFunction.cat([dt_phi_x + adv_x, dt_phi_v + adv_v, poisson_res]) class ICFlowResidual(InitialResidual): """CI du flot : T^{-1}(x,v,0) = (x,v) (identité à t=0).""" 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, inv_phi, densite_f, phi_e, rho, u, T, ) -> PARAM_FUNC_TYPE: return inv_phi.set_t_0(self.time_domain[0]) class VlasovPoissonModel(AbstractPhysicalModel): def __init__( self, main_domain: VolumetricDomain, time_domain: tuple[float, float], f_ic_rhs: NDARRAYS_FUNC_TYPE | None = None, ): super().__init__(main_domain=main_domain, time_domain=time_domain) label = main_domain.get_label() self.physical_residuals: PHYSICAL_RESIDUALS_TYPE = { label: VlasovPoissonResidual(domain=main_domain, time_domain=time_domain), "ic " + label: ICFlowResidual( domain=main_domain, time_domain=(time_domain[0],), f_rhs=f_ic_rhs, ), } def renew_time_domain(self, new_time_domain: tuple[float, float]): self.time_domain = new_time_domain for res in self.physical_residuals.values(): res.time_domain = new_time_domain # ───────────────────────────────────────────────────────────────────────────── # Configuration # ───────────────────────────────────────────────────────────────────────────── N_COLLOC = 5_000 N_IC_COLLOC = 2_000 N_EPOCHS = [200, 200] # nb d'époques par flot (int ou list[int]) dim_x = 1 # dimension de la position dim_v = 1 # dimension de la vitesse params_dim = 0 nb_models = 2 # nombre de flots temporels Tf = 2.0 Deltat = Tf / nb_models time_intervals = [ (0.0, 1.0), (1.0, 2.0), ] # [(Deltat * i, Deltat * (i + 1)) for i in range(nb_models)] key = jax.random.PRNGKey(42) # ───────────────────────────────────────────────────────────────────────────── # Domaines # ───────────────────────────────────────────────────────────────────────────── domain_x = Segment1D( jnp.array([[0.0, float(X_L)]]), is_main_domain=True, ) domain_v = Segment1D(jnp.array([[-V_MAX, V_MAX]])) # ───────────────────────────────────────────────────────────────────────────── # Réseaux de neurones # ───────────────────────────────────────────────────────────────────────────── keys = jax.random.split(key, nb_models + 1) # Flot arrière φ : (x, v, t) → (x_0, v_0) sizes = [[30, 30], [30, 30]] flow_models = [ MLP( in_size=dim_x + dim_v + 1, # x + v + t hidden_sizes=sizes[i], out_size=dim_x + dim_v, # (x_0, v_0) activation="tanh", key=keys[i], embedding="periodic", periods=(float(X_L),), # période de x embedding_axes=(0,), # seul x est périodique (axe 0), v et t restent linéaires ) for i in range(nb_models) ] # Alternative MLP périodique en x : embedding (cos, sin) sur l'axe x # flow_models = [ # MLP( # in_size=dim_x + dim_v + 1, # x + v + t (entrée brute) # hidden_sizes=[26, 26, 26], # out_size=dim_x + dim_v, # activation="silu", # key=keys[i], # embedding="periodic", # periods=(float(X_L),), # période de x # embedding_axes=(0,), # seul x est périodique (axe 0) # ) # for i in range(nb_models) # ] # Alternative symplectique : |det J| = 1 exact → pas de jacrev dans compute_density # Utiliser avec add_id=False dans DensityFlowFieldsApproximationSpace. # from scimba_jax.nonlinear_approximation.networks.structure_preserving_nets.symplectic_nets import ( # PeriodicGSymplecticNet, # ) # flow_models = [ # PeriodicGSymplecticNet( # size=dim_x + dim_v, # dimension espace des phases (2) # conditional_size=1, # t # width=32, # nb_layers=8, # key=keys[i], # period=jnp.array([float(X_L)]), # périodicité en x (1ère coord) # activation="tanh", # h=0.025, # ) # for i in range(nb_models) # ] # Potentiel électrique φ_E : (x, t) → φ_E(x,t), périodique en x # PeriodicEmbedding remplace x par (cos(2πx/L), sin(2πx/L)) → out_size = 2+1 = 3 phi_e_model = MLP( in_size=dim_x + 1, # x + t (entrée brute) hidden_sizes=[30, 30], out_size=1, activation="tanh", key=keys[-1], embedding="periodic", periods=(float(X_L),), # période de x embedding_axes=(0,), # seul x est périodique (axe 0), t reste linéaire ) # ───────────────────────────────────────────────────────────────────────────── # Quadrature en vitesse pour le calcul des moments # ───────────────────────────────────────────────────────────────────────────── # UnitSquareTensorized génère des points dans [0,1] ; on remet sur [-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 # ───────────────────────────────────────────────────────────────────────────── # Espace d'approximation # ───────────────────────────────────────────────────────────────────────────── space = DensityFlowFieldsApproximationSpace( dim=dim_x, velocity_dim=dim_v, params_dim=params_dim, models=flow_models, list_models_fields=[(phi_e_model, "scalar", 1)], # (modèle, type, taille sortie) add_id=True, # φ(x,v,t) = (x,v) + réseau (proche de l'identité) time_intervals=time_intervals, time_continuous=True, parametric_dependence=False, initial_density=f0, moment_computation=True, quadrature_velocity=quad_velocity, scan_moments=False, ) # ───────────────────────────────────────────────────────────────────────────── # Sampler : domaine en espace × domaine en vitesse × temps # ───────────────────────────────────────────────────────────────────────────── initial_t_domain = (0.0, time_intervals[0][1]) sampler = TensorizedSampler( [ DomainSampler(domain_x), # "x" : position 1D UniformVelocitySamplerOnCuboid(domain_v), # "v" : vitesse 1D UniformTimeSampler(initial_t_domain), # "t" : temps ], bc=False, ic=True, model_type="x_v_t", ) # ───────────────────────────────────────────────────────────────────────────── # Modèle physique # ───────────────────────────────────────────────────────────────────────────── model = VlasovPoissonModel( main_domain=domain_x, time_domain=initial_t_domain, f_ic_rhs=ic_flow, ) weights = { "interior": [2.0, 2.0, 1.0], # φ_x Vlasov, φ_v Vlasov, Poisson "ic interior": [5.0, 5.0], # CI φ_x, CI φ_v } # ───────────────────────────────────────────────────────────────────────────── # Boucle d'entraînement (un flot par intervalle de temps) # ───────────────────────────────────────────────────────────────────────────── list_of_pinns = [] pinn_total_time = 0.0 for i in range(nb_models): space.idx_current_flow = i # space.warm_start_from_previous_flow() domain_t = time_intervals[i] print(f"\n=== Flot {i + 1}/{nb_models} — t ∈ {domain_t} ===") sampler.renew_sampler(UniformTimeSampler(domain_t), "t") model.renew_time_domain(domain_t) pinn = Projector( model, space, sampler, optimizer="ENG", weights=weights, one_loss_per_residual=True, matrix_regularization=1.0e-6, ) start = timeit.default_timer() n_epochs_i = N_EPOCHS[i] if isinstance(N_EPOCHS, list) else N_EPOCHS key, pinn = pinn.project( key, space, n_epochs=n_epochs_i, n_colloc=N_COLLOC, n_ic_colloc=N_IC_COLLOC ) end = timeit.default_timer() list_of_pinns.append(copy.deepcopy(pinn)) space = pinn.space elapsed = end - start pinn_total_time += elapsed print(f"\n=== PINN : temps total d'entraînement : {pinn_total_time:.1f}s ===") # ───────────────────────────────────────────────────────────────────────────── # Visualisation # ───────────────────────────────────────────────────────────────────────────── def _maxwellian_np(vv): import numpy as np return np.exp(-0.5 * vv**2) / np.sqrt(2.0 * np.pi) def plot_phase_space(pinn, t_vals=(0.0, 5.0, 10.0, 20.0), n_visu=128): """Visualise δf = f−M(v), ρ−1 et φ_E pour plusieurs instants. Indices dans vals : 0,1→flot 2→f 3→φ_E 4→ρ 5→u 6→T """ import numpy as np x_lin = np.linspace(0.0, float(X_L), n_visu) v_lin = np.linspace(-V_MAX, V_MAX, n_visu) xx, vv = np.meshgrid(x_lin, v_lin) x_flat = jnp.array(xx.ravel()[:, None]) v_flat = jnp.array(vv.ravel()[:, None]) M_v = _maxwellian_np(vv) # Amplitude de δf à t=0 pour normalisation cohérente t0 = jnp.zeros_like(x_flat) f0 = np.array(pinn.evaluate(x_flat, v_flat, t0))[:, 2].reshape(n_visu, n_visu) lim_df = max(np.abs(f0 - M_v).max(), 1e-12) nt = len(t_vals) fig, axes = plt.subplots(4, nt, figsize=(4.5 * nt, 13)) for col, t_val in enumerate(t_vals): t_flat = jnp.full_like(x_flat, t_val) vals = np.array(pinn.evaluate(x_flat, v_flat, t_flat)) f_vals = vals[:, 2].reshape(n_visu, n_visu) df = f_vals - M_v # ── f(x,v,t) ───────────────────────────────────────────────────────── f0_vals = np.array(pinn.evaluate(x_flat, v_flat, t0))[:, 2].reshape( n_visu, n_visu ) lev_f = np.linspace(0.0, f0_vals.max(), 40) im0 = axes[0, col].contourf( xx, vv, f_vals, levels=lev_f, cmap="inferno", extend="both" ) axes[0, col].contour( xx, vv, f_vals, levels=10, colors="white", linewidths=0.5, alpha=0.6 ) axes[0, col].set_title(f"f(x,v, t={t_val:.1f})") axes[0, col].set_xlabel("x") axes[0, col].set_ylabel("v") plt.colorbar(im0, ax=axes[0, col]) # ── δf = f − M(v) ──────────────────────────────────────────────────── lev = np.linspace(-lim_df, lim_df, 41) im0b = axes[1, col].contourf( xx, vv, df, levels=lev, cmap="RdBu_r", extend="both" ) axes[1, col].contour( xx, vv, df, levels=10, colors="k", linewidths=0.3, alpha=0.5 ) axes[1, col].set_title(f"δf(x,v, t={t_val:.1f})") axes[1, col].set_xlabel("x") axes[1, col].set_ylabel("v") plt.colorbar(im0b, ax=axes[1, col], format="%.2e") # ── ρ−1 et φ_E sur grille x ────────────────────────────────────────── x1d = jnp.array(x_lin[:, None]) t1d = jnp.full_like(x1d, t_val) vals1d = np.array(pinn.evaluate(x1d, jnp.zeros_like(x1d), t1d)) rho_m1 = vals1d[:, 4] - 1.0 # ρ−1 phi_E = vals1d[:, 3] # φ_E axes[2, col].plot(x_lin, rho_m1, color="steelblue") axes[2, col].axhline(0.0, color="gray", linestyle="--", linewidth=0.8) axes[2, col].set_title(f"ρ−1 (x, t={t_val:.1f})") axes[2, col].set_xlabel("x") axes[3, col].plot(x_lin, phi_E, color="tomato") axes[3, col].axhline(0.0, color="gray", linestyle="--", linewidth=0.8) axes[3, col].set_title(f"φ_E(x, t={t_val:.1f})") axes[3, col].set_xlabel("x") plt.tight_layout() plt.show() def plot_landau_damping(pinn, n_t=100, n_x=64): """log|E_max|(t) vs taux théorique Landau (k=0.5 : γ≈−0.1533).""" import numpy as np x1d = jnp.linspace(0.0, float(X_L), n_x)[:, None] v0 = jnp.zeros_like(x1d) t_vals = np.linspace(0.0, float(Tf), n_t) E_max = [] for t_val in t_vals: t1d = jnp.full_like(x1d, t_val) vals = np.array(pinn.evaluate(x1d, v0, t1d)) phi_E = vals[:, 3] E_x = -np.gradient(phi_E, float(X_L) / n_x, edge_order=2) E_max.append(np.abs(E_x).max()) gamma_th = -0.1533 E0 = E_max[0] if E_max[0] > 0 else 1e-10 plt.figure(figsize=(8, 4)) plt.semilogy(t_vals, E_max, color="tomato", lw=1.5, label="|E|_max (PINN)") plt.semilogy( t_vals, E0 * np.exp(gamma_th * t_vals), "k--", lw=1, label=f"théorie γ={gamma_th}", ) plt.xlabel("t") plt.ylabel("|E|_max") plt.title("Amortissement de Landau — log|E|(t)") plt.legend() plt.tight_layout() plt.show() def plot_flow_components(pinn, t_vals=None, n_visu=64): """Plot δφ_x = φ_x−x et δφ_v = φ_v−v (perturbation du flot arrière composé).""" import numpy as np if t_vals is None: t_vals = [0.0, Tf / 4, Tf / 2, 3 * Tf / 4, Tf] space = pinn.space _eval_flow = jax.jit( jax.vmap( lambda xvt: space.compute_flow_backward_compose_continuous_inference( space, xvt ) ) ) 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) xx, vv = np.meshgrid(x_visu, v_visu) x_flat = jnp.array(xx.ravel()[:, None]) v_flat = jnp.array(vv.ravel()[:, None]) xv_flat = jnp.concatenate([x_flat, v_flat], axis=-1) nt = len(t_vals) fig, axes = plt.subplots(2, nt, figsize=(4.0 * nt, 7), squeeze=False) fig.suptitle("Flot arrière composé — δφ_x = φ_x−x | δφ_v = φ_v−v", fontsize=13) for col, t_val in enumerate(t_vals): t_col = jnp.full((n_visu * n_visu, 1), t_val) flow = np.array( _eval_flow(jnp.concatenate([xv_flat, t_col], axis=-1)) ) # (n², 2) dphi_x = flow[:, 0].reshape(n_visu, n_visu) - xx dphi_v = flow[:, 1].reshape(n_visu, n_visu) - vv 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, vv, data, levels=40, cmap="RdBu_r", vmin=-vmax, vmax=vmax ) ax.contour(xx, vv, 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() plt.show() def compute_electric_energy_from_rho(rho, dx): """Résout Poisson par FFT et retourne (1/2)∫E²dx depuis ρ(x).""" Nx = rho.shape[0] kx = jnp.fft.fftfreq(Nx) * Nx * 2.0 * jnp.pi / float(X_L) rho_hat = jnp.fft.fft(rho - 1.0) E_hat = jnp.where(kx != 0, -1j * rho_hat / kx, 0.0 + 0.0j) return 0.5 * jnp.sum(jnp.fft.ifft(E_hat).real ** 2) * dx def plot_energy_decay( pinn, n_t=80, Nx_fine=NX_SL, Nv_fine=NV_SL, dt_fine=DT_SL, Nx_coarse=64, Nv_coarse=64, dt_coarse=0.05, n_pinn_x=128, n_pinn_v=128, ): """Décroissance de l'énergie électrique : SL fin, SL coarse, PINN-from-flow.""" import numpy as np t_vals = list(np.linspace(0.0, float(Tf), n_t)) def _sl_energies(Nx, Nv, dt, label): print(f"SL {label} ({Nx}x{Nv}, dt={dt})…") results, t_actual, _, _ = run_sl_vlasov_poisson(t_vals, Nx=Nx, Nv=Nv, dt=dt) dx = float(X_L) / Nx dv = 2.0 * V_MAX / Nv energies = [] for f_sl, _, _ in results: rho = jnp.array(f_sl).sum(axis=1) * dv # f shape (Nx,Nv) → sum over v energies.append(float(compute_electric_energy_from_rho(rho, dx))) return np.array(t_actual), np.array(energies) t_fine, E_fine = _sl_energies(Nx_fine, Nv_fine, dt_fine, "fin") t_coarse, E_coarse = _sl_energies(Nx_coarse, Nv_coarse, dt_coarse, "coarse") print("PINN energy from flow…") space = pinn.space _eval_f = jax.jit(jax.vmap(lambda xvt: space.compute_density(space, xvt))) x_p = np.linspace(0.0, float(X_L), n_pinn_x, endpoint=False) v_p = np.linspace(-V_MAX, V_MAX, n_pinn_v, endpoint=False) dv_p = 2.0 * V_MAX / n_pinn_v dx_p = float(X_L) / n_pinn_x xx_p, vv_p = np.meshgrid(x_p, v_p) x_flat = jnp.array(xx_p.ravel()[:, None]) v_flat = jnp.array(vv_p.ravel()[:, None]) xv_flat = jnp.concatenate([x_flat, v_flat], axis=-1) E_pinn = [] for t_val in t_vals: t_col = jnp.full((n_pinn_x * n_pinn_v, 1), t_val) f_vals = np.array(_eval_f(jnp.concatenate([xv_flat, t_col], axis=-1))) rho = jnp.array(f_vals.reshape(n_pinn_v, n_pinn_x).sum(axis=0) * dv_p) E_pinn.append(float(compute_electric_energy_from_rho(rho, dx_p))) E_pinn = np.array(E_pinn) t_pinn = np.array(t_vals) gamma_th = -0.1533 E0_ref = E_fine[0] if E_fine[0] > 0 else 1e-10 plt.figure(figsize=(9, 4)) plt.semilogy(t_fine, E_fine, color="steelblue", lw=1.5, label="SL fin") plt.semilogy(t_coarse, E_coarse, ":", color="steelblue", lw=1.5, label="SL coarse") plt.semilogy(t_pinn, E_pinn, color="tomato", lw=1.5, label="PINN (from flow)") plt.semilogy( t_pinn, E0_ref * np.exp(2 * gamma_th * t_pinn), "k--", lw=1, label=f"théorie exp(2γt), γ={gamma_th}", ) plt.xlabel("t") plt.ylabel(r"$\frac{1}{2}\int E^2 dx$") plt.title("Décroissance de l'énergie électrique") plt.legend() plt.tight_layout() plt.show() def plot_comparison_sl_pinn( pinn, t_vals=None, Nx=NX_SL, Nv=NV_SL, dt=DT_SL, n_visu=128, Nx_coarse=128, Nv_coarse=128, dt_coarse=0.01, ): """Compare f, δf, ρ et φ entre PINN et SL à plusieurs instants. Plot 1 : grille 4×nt — lignes : f PINN | δf PINN | f SL | δf SL. Plot 2 : grille 2×nt — lignes : ρ−1 | φ (PINN vs SL sur le même axe). n_visu : résolution PINN (n_visu × n_visu). Nx_coarse : résolution SL coarse (ajouté aux coupes 1D). """ import numpy as np if t_vals is None: t_vals = [0.0, Tf / 4, Tf / 2, 3 * Tf / 4, Tf] # ── SL : une passe unique, sauvegarde aux checkpoints ──────────────────── sl_start = timeit.default_timer() sl_list, t_actual, x_sl, v_sl = run_sl_vlasov_poisson(t_vals, Nx=Nx, Nv=Nv, dt=dt) sl_elapsed = timeit.default_timer() - sl_start print(f"\n=== SL : temps de simulation : {sl_elapsed:.1f}s ===") x_sl = np.array(x_sl) v_sl = np.array(v_sl) # Grille PINN réduite (n_visu × n_visu) 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) # (n_visu, n_visu) x_flat = jnp.array(xx_pinn.ravel()[:, None]) v_flat = jnp.array(vv_pinn.ravel()[:, None]) M_v_pinn = np.exp(-0.5 * vv_pinn**2) / np.sqrt(2.0 * np.pi) # Grille SL pour les plots 2D xx_sl, vv_sl = np.meshgrid(x_sl, v_sl) # (Nv, Nx) M_v_sl = np.exp(-0.5 * vv_sl**2) / np.sqrt(2.0 * np.pi) nt = len(t_actual) f_pinns, df_pinns, f_sls, df_sls = [], [], [], [] rho_pinns, phi_pinns, rho_sls, phi_sls_list = [], [], [], [] # ── SL coarse aux mêmes checkpoints ───────────────────────────────────── coarse_list, _, x_sl_coarse, v_sl_coarse = run_sl_vlasov_poisson( t_vals, Nx=Nx_coarse, Nv=Nv_coarse, dt=dt_coarse ) x_sl_coarse = np.array(x_sl_coarse) v_sl_coarse = np.array(v_sl_coarse) f_coarses = [] # Évaluation directe : f via compute_density, φ_E via le MLP — sans moments space = pinn.space _eval_f = jax.jit(jax.vmap(lambda xvt: space.compute_density(space, xvt))) _eval_phi = jax.jit(jax.vmap(lambda xt: space.fields_models[0](xt)[0])) x_visu_jax = jnp.array(x_visu[:, None]) xv_flat = jnp.concatenate( [x_flat, v_flat], axis=-1 ) # (N, 2), t ajouté dans la boucle for idx, t_val in enumerate(t_actual): # ── PINN f sur la grille 2D (un seul jacrev par point, pas de moments) ── t_col = jnp.full((n_visu * n_visu, 1), t_val) f_vals = np.array(_eval_f(jnp.concatenate([xv_flat, t_col], axis=-1))) f_pinn = f_vals.reshape(n_visu, n_visu) f_pinns.append(f_pinn) df_pinns.append(f_pinn - M_v_pinn) # ρ par intégration numérique de f sur la grille v (rectangles) rho_pinns.append(f_pinn.sum(axis=0) * dv_visu) # (n_visu,) # φ_E via le MLP directement sur la grille x 1D t_1d = jnp.full((n_visu, 1), t_val) phi_vals = np.array(_eval_phi(jnp.concatenate([x_visu_jax, t_1d], axis=-1))) phi_pinn_i = EPS * phi_vals phi_pinn_i -= phi_pinn_i.mean() phi_pinns.append(phi_pinn_i) # ── SL coarse ───────────────────────────────────────────────────────── f_c, _, _ = coarse_list[idx] f_coarses.append(np.array(f_c).T) # (Nv_coarse, Nx_coarse) # ── SL fin ─────────────────────────────────────────────────────────────── f_sl_i, rho_sl_i, phi_sl_i = sl_list[idx] f_sl_arr = np.array(f_sl_i).T # (Nx,Nv).T → (Nv,Nx) f_sls.append(f_sl_arr) df_sls.append(f_sl_arr - M_v_sl) rho_sls.append(np.array(rho_sl_i)) phi_sl_c = np.array(phi_sl_i) phi_sl_c -= phi_sl_c.mean() phi_sls_list.append(phi_sl_c) # ── Niveaux communs ─────────────────────────────────────────────────────── 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) # ── Plot 1 : f et δf ───────────────────────────────────────────────────── fig1, axes1 = plt.subplots(4, nt, figsize=(4.5 * nt, 14), squeeze=False) fig1.suptitle("f et δf = f−M(v) — PINN 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), (df_sls, "RdBu_r", lev_df, "%.2e", xx_sl, vv_sl), ] 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[col]:.3f}") axes1[row, 0].set_ylabel("v") for row, (label, color) in enumerate( [ ("PINN", "steelblue"), ("PINN", "steelblue"), ("SL", "tomato"), ("SL", "tomato"), ] ): 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) plt.show() # ── Plot 2 : ρ−1 et φ ──────────────────────────────────────────────────── fig2, axes2 = plt.subplots(2, nt, figsize=(4.5 * nt, 7), squeeze=False) fig2.suptitle("ρ−1 et φ — PINN vs SL", fontsize=13) for col in range(nt): t_val = t_actual[col] ax = axes2[0, col] ax.plot(x_sl, rho_sls[col] - 1.0, color="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", linestyle=":", lw=0.8) ax.set_title(f"ρ−1 t = {t_val:.3f}") ax.set_xlabel("x") if col == 0: ax.legend() ax = axes2[1, col] ax.plot(x_sl, phi_sls_list[col], color="steelblue", lw=1.5, label="SL") ax.plot(x_visu, phi_pinns[col], "--", color="tomato", lw=1.5, label="PINN") ax.axhline(0.0, color="gray", linestyle=":", lw=0.8) ax.set_title(f"φ t = {t_val:.3f}") ax.set_xlabel("x") if col == 0: ax.legend() plt.tight_layout() plt.show() # ── Plots 3-6 : coupes 1D ──────────────────────────────────────────────── def _cut_v(x_val, label): j_p = int(np.argmin(np.abs(x_visu - x_val))) j_c = int(np.argmin(np.abs(x_sl_coarse - 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"Coupe {label} — PINN / SL fin (--) / SL coarse (:)", 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") ax.plot( v_sl, f_sls[col][:, j_s], "--", color="steelblue", lw=1.5, label="SL fin", ) ax.plot( v_sl_coarse, f_coarses[col][:, j_c], ":", color="steelblue", lw=1.2, label="SL coarse", ) ax.set_title(f"t = {t_actual[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() plt.show() def _cut_x(v_val, label): i_p = int(np.argmin(np.abs(v_visu - v_val))) i_c = int(np.argmin(np.abs(v_sl_coarse - 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"Coupe {label} — PINN / SL fin (--) / SL coarse (:)", 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") ax.plot( x_sl, f_sls[col][i_s, :], "--", color="steelblue", lw=1.5, label="SL fin", ) ax.plot( x_sl_coarse, f_coarses[col][i_c, :], ":", color="steelblue", lw=1.2, label="SL coarse", ) ax.set_title(f"t = {t_actual[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() plt.show() _cut_v(np.pi * 1.0, "x=π") _cut_v(np.pi * 2.0, "x=2π") _cut_v(np.pi * 3.0, "x=3π") _cut_x(0.0, "v=0") # %% if plot := True: for i, pinn in enumerate(list_of_pinns): pinn.space.idx_current_flow = i for label in pinn.losses.losses_history: plt.semilogy(pinn.losses.losses_history[label], label=label) plt.title(f"Loss — flot {i}") plt.legend() plt.show() pinn_last = list_of_pinns[-1] pinn_last.space.idx_current_flow = len(list_of_pinns) - 1 pinn_last.space.inference_mode = True plot_flow_components(pinn_last, t_vals=[0.0, Tf / 4, Tf / 2, 3 * Tf / 4, Tf]) plot_energy_decay(pinn_last) plot_comparison_sl_pinn(pinn_last, t_vals=[0.0, Tf / 4, Tf / 2, 3 * Tf / 4, Tf])