# %% """Vlasov-Poisson 1D bi-niveau — 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) Approche bi-niveau : — Flot arrière φ = (φ_x, φ_v) : résidu Vlasov seul (pas de Poisson) — φ_E(x,t) : résolu par un PINN Poisson séparé (field solver) ∂_xx φ_E + (ρ−1)/ε = 0, ρ = ∫ f dv (source issue du flot arrière) — Les poids θ_p du PINN Poisson sont calculés implicitement depuis θ_f (flot) via DensityFlowBiLevelVlasovPoisson.moment_to_fields_solver """ # %% 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.approximation_spaces import ( ApproximationSpace, ) from scimba_jax.nonlinear_approximation.approximation_spaces.densityflowbilevel_approximation_spaces import ( DensityFlowBiLevelApproximationSpace, ) 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 = "two_stream" K = 0.5 X_L = 2.0 * jnp.pi / K if CASE == "landau_damping": EPS = 0.1 V_MAX = 6.0 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"): 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, *_): return jnp.concatenate([x, v], axis=-1) # ───────────────────────────────────────────────────────────────────────────── # Solveur Semi-Lagrangien Vlasov-Poisson 1D (référence numérique) # ───────────────────────────────────────────────────────────────────────────── 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): 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") xv = jnp.stack([xx, vv], axis=-1) f = f0(xv, case=CASE) 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) 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**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) 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) 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) 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 : {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 Vlasov (flot arrière uniquement — sans Poisson) # ───────────────────────────────────────────────────────────────────────────── class VlasovResidualBilevel(InteriorResidual): """Résidu Vlasov seul (taille 2). Variables reçues dans l'ordre de create_variables du bilevel space : 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) [non utilisé] phi_e — potentiel électrique φ_E(x,t) du field solver [scalaire] Équation : ∂_t φ_i + v ∂_x φ_i − ε ∂_x φ_E ∂_v φ_i = 0 (i = x, v) """ def __init__(self, domain: VolumetricDomain, time_domain, f_rhs=None): super().__init__( domain=domain, size=2, model_type="x_v_t_dofsl", f_rhs=f_rhs, time_domain=time_domain, ) def construct_residual( self, inv_phi: PARAM_FUNC_TYPE, _densite_f: PARAM_FUNC_TYPE, phi_e: PARAM_FUNC_TYPE, ) -> PARAM_FUNC_TYPE: phi_x, phi_v = inv_phi.components() dt_phi_x = phi_x.d_t() dt_phi_v = phi_v.d_t() dphi_e_dx = phi_e.gradient("x") E_val = -EPS * dphi_e_dx def _advection(phi_comp): dphi_dx = phi_comp.gradient("x") dphi_dv = phi_comp.gradient("v") def v_func(v): return v 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) return ParamVecFunction.cat([dt_phi_x + adv_x, dt_phi_v + adv_v]) class ICFlowResidualBilevel(InitialResidual): """CI du flot : T^{-1}(x,v,0) = (x,v).""" 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_dofsl", f_rhs=f_rhs, ) def construct_residual( self, inv_phi: PARAM_FUNC_TYPE, _densite_f: PARAM_FUNC_TYPE, phi_e: PARAM_FUNC_TYPE, ) -> PARAM_FUNC_TYPE: res = inv_phi.set_t_0(self.time_domain[0]) # Ce résidu ne dépend pas de phi_e (donc pas de θ_p*), mais son f_type doit # quand même porter un "dofsl" final : gram_matrix_function/assembly_post_sampling # appendent uniformément int_values=(theta_p_flat,) à TOUS les résidus # (new_args = sample_dict[key] + int_values), donc l'arité de res doit # matcher model_type="x_v_dofsl" (cf. get_vmap_axes). theta_p_flat est # simplement ignoré par wrapper → jacobien nul vis-à-vis de θ_p*. res_vars = res.f_type.split("_") promoted_f_type = res.f_type + "_dofsl" promoted_dims = {k: res.dims[k] for k in res_vars} | { "dofsl": phi_e.dims["dofsl"] } def wrapper(space, *args): return res(space, *args[: len(res_vars)]) return ParamVecFunction( res.size, promoted_dims, wrapper, f_type=promoted_f_type ) class VlasovModelBilevel(AbstractPhysicalModel): """Modèle Vlasov seul — Poisson résolu séparément par le field solver.""" 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: VlasovResidualBilevel(domain=main_domain, time_domain=time_domain), "ic " + label: ICFlowResidualBilevel( 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 # ───────────────────────────────────────────────────────────────────────────── # Résidu et modèle Poisson (standard scimba) # ───────────────────────────────────────────────────────────────────────────── class PoissonResidual(InteriorResidual): """Résidu -∆φ_E = f_rhs, f_rhs = (ρ-1)/ε (source fournie via f_rhs).""" def __init__( self, domain: VolumetricDomain, time_domain: tuple[float, float], f_rhs: NDARRAYS_FUNC_TYPE | None = None, ): super().__init__( domain=domain, size=1, model_type="x_t", f_rhs=f_rhs, time_domain=time_domain, ) def construct_residual(self, phi_e: PARAM_FUNC_TYPE) -> PARAM_FUNC_TYPE: return -phi_e.laplacian("x") class PoissonModel(AbstractPhysicalModel): """PINN Poisson espace-temps : -∆φ_E = (ρ-1)/ε.""" def __init__( self, main_domain: VolumetricDomain, time_domain: tuple[float, float], f_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: PoissonResidual( domain=main_domain, time_domain=time_domain, f_rhs=f_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 # ───────────────────────────────────────────────────────────────────────────── # Espace bi-niveau Vlasov-Poisson # ───────────────────────────────────────────────────────────────────────────── class DensityFlowBiLevelVlasovPoisson(DensityFlowBiLevelApproximationSpace): """Espace bi-niveau pour Vlasov-Poisson 1D. Implémente moment_to_fields_solver : câble ρ(x,t) = ∫f(x,v,t)dv comme terme source du résidu Poisson du field solver. ρ est calculé via la quadrature de Gauss en v (quadrature_velocity). Le field solver (PINN Poisson) est ensuite résolu pour obtenir φ_E(x,t;θ_p*(θ_f)). """ def moment_to_fields_solver(self, _var_density=None): ks = self def rhs_fn(x, t): # -∆φ_E = f_rhs ⟹ f_rhs = (ρ-1)/ε (convention : ∆φ_E + (ρ-1)/ε = 0) # ks est tracé par JAX → compute_moments(ks, …) porte le gradient en θ_f return (ks.compute_moments(ks, x, t)[0] - 1.0) / EPS # On ne reconstruit pas PoissonModel (son __init__ appelle get_all_boundaries() # qui fait .item() sur un tableau JAX tracé → ConcretizationTypeError). # On met seulement à jour le résidu existant (PoissonResidual.__init__ est safe) # puis on reconstruit le Projector pour obtenir un Evaluator frais avec le nouveau f_rhs. # NOTE : on NE réassigne PAS ks.field_solver — field_solver est en aux_data # (statique) et doit garder une identité stable entre l'entrée et la sortie # de _field_solve pour jax.custom_vjp (cf. PyTreeDef structure check). old = ks.field_solver label = old.model.main_domain.get_label() old.model.physical_residuals[label] = PoissonResidual( domain=old.model.main_domain, time_domain=old.model.time_domain, f_rhs=rhs_fn, ) return Projector(old.model, old.space, old.sampler) # ───────────────────────────────────────────────────────────────────────────── # Configuration # ───────────────────────────────────────────────────────────────────────────── N_COLLOC = 5_000 N_IC_COLLOC = 2_000 N_EPOCHS_VLASOV = [400] N_EPOCHS_POISSON = 20 # époques du PINN Poisson à chaque solve bi-niveau dim_x = 1 dim_v = 1 params_dim = 0 nb_models = 1 Tf = 1.0 time_intervals = [ (0.0, 2.0), ] 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 + 2) # Flot arrière φ : (x, v, t) → (x_0, v_0) sizes = [[20, 20]] flow_models = [ MLP( in_size=dim_x + dim_v + 1, hidden_sizes=sizes[i], 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) ] # Potentiel électrique φ_E : (x, t) → φ_E(x,t), périodique en x phi_e_model = MLP( in_size=dim_x + 1, hidden_sizes=[16, 16], out_size=1, activation="tanh", key=keys[nb_models], embedding="periodic", periods=(float(X_L),), embedding_axes=(0,), ) # ───────────────────────────────────────────────────────────────────────────── # Quadrature en vitesse pour ρ = ∫f dv dans moment_to_fields_solver # ───────────────────────────────────────────────────────────────────────────── 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 # ───────────────────────────────────────────────────────────────────────────── # PINN Poisson (field solver) # ───────────────────────────────────────────────────────────────────────────── initial_t_domain = (0.0, time_intervals[0][1]) poisson_sampler = TensorizedSampler( [ DomainSampler(domain_x), UniformTimeSampler(initial_t_domain), ], bc=False, ic=False, model_type="x_t", ) # Espace standard ApproximationSpace — create_variables() construit phi_e automatiquement poisson_space = ApproximationSpace( {"x": dim_x, "t": 1}, [(phi_e_model, "scalar", None)], model_type="x_t", ) # Modèle Poisson initial (f_rhs=None : sera câblé dans moment_to_fields_solver) poisson_model = PoissonModel(domain_x, time_domain=initial_t_domain) poisson_pinn = Projector(poisson_model, poisson_space, poisson_sampler) # ───────────────────────────────────────────────────────────────────────────── # Espace bi-niveau # ───────────────────────────────────────────────────────────────────────────── bilevel_space = DensityFlowBiLevelVlasovPoisson( dim=dim_x, velocity_dim=dim_v, params_dim=params_dim, models=flow_models, add_id=True, time_intervals=time_intervals, time_continuous=True, parametric_dependence=False, initial_density=f0, moment_computation=False, quadrature_velocity=quad_velocity, scan_moments=False, field_solver=poisson_pinn, field_epochs=N_EPOCHS_POISSON, ) # ───────────────────────────────────────────────────────────────────────────── # Sampler Vlasov # ───────────────────────────────────────────────────────────────────────────── vlasov_sampler = TensorizedSampler( [ DomainSampler(domain_x), UniformVelocitySamplerOnCuboid(domain_v), UniformTimeSampler(initial_t_domain), ], bc=False, ic=True, model_type="x_v_t", ) # ───────────────────────────────────────────────────────────────────────────── # Modèle Vlasov bi-niveau # ───────────────────────────────────────────────────────────────────────────── vlasov_model = VlasovModelBilevel( main_domain=domain_x, time_domain=initial_t_domain, f_ic_rhs=ic_flow, ) weights_vlasov = { "interior": [2.0, 2.0], # φ_x Vlasov, φ_v Vlasov "ic interior": [5.0, 5.0], } # ───────────────────────────────────────────────────────────────────────────── # Boucle d'entraînement bi-niveau # ───────────────────────────────────────────────────────────────────────────── list_of_pinns = [] pinn_total_time = 0.0 for i in range(nb_models): bilevel_space.idx_current_flow = i domain_t = time_intervals[i] print(f"\n=== Flot {i + 1}/{nb_models} — t ∈ {domain_t} ===") # Mettre à jour les domaines temporels vlasov_sampler.renew_sampler(UniformTimeSampler(domain_t), "t") vlasov_model.renew_time_domain(domain_t) poisson_sampler.renew_sampler(UniformTimeSampler(domain_t), "t") poisson_model.renew_time_domain(domain_t) # Reconstruire le field solver Poisson pour le nouvel intervalle temporel # (f_rhs=None ici ; moment_to_fields_solver le câblera à chaque solve interne) poisson_pinn = Projector(poisson_model, poisson_space, poisson_sampler) bilevel_space.field_solver = poisson_pinn # Projector Vlasov (niveau externe du bi-niveau) vlasov_pinn = Projector( vlasov_model, bilevel_space, vlasov_sampler, optimizer="SS-BFGS", weights=weights_vlasov, learning_rate=1e-2, # matrix_regularisation=1e-3, ) start = timeit.default_timer() n_epochs_i = ( N_EPOCHS_VLASOV[i] if isinstance(N_EPOCHS_VLASOV, list) else N_EPOCHS_VLASOV ) key, vlasov_pinn = vlasov_pinn.project( key, bilevel_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(vlasov_pinn)) pinn_total_time += end - start print(f"\n=== PINN bi-niveau : temps total d'entraînement : {pinn_total_time:.1f}s ===") # ───────────────────────────────────────────────────────────────────────────── # Visualisation # ───────────────────────────────────────────────────────────────────────────── # Variables dans bilevel_space.create_variables() : # vals[:, 0:2] → flot arrière (φ_x, φ_v) # vals[:, 2] → densité f(x,v,t) # vals[:, 3] → potentiel φ_E(x,t) du field solver 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_bilevel(pinn, t_vals=(0.0, 3.0, 6.0, 9.0), n_visu=128): """Visualise f et φ_E pour plusieurs instants (version bi-niveau). Indices : 0,1 → flot | 2 → f | 3 → φ_E """ 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) t0 = jnp.zeros_like(x_flat) f0_arr = np.array(pinn.evaluate(x_flat, v_flat, t0))[:, 2].reshape(n_visu, n_visu) lim_df = max(np.abs(f0_arr - M_v).max(), 1e-12) nt = len(t_vals) fig, axes = plt.subplots(3, nt, figsize=(4.5 * nt, 10)) 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 lev_f = np.linspace(0.0, f0_arr.max(), 40) im0 = axes[0, col].contourf( xx, vv, f_vals, levels=lev_f, cmap="inferno", extend="both" ) 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]) lev = np.linspace(-lim_df, lim_df, 41) im1 = axes[1, col].contourf( xx, vv, df, levels=lev, cmap="RdBu_r", extend="both" ) axes[1, col].set_title(f"δf(x,v, t={t_val:.1f})") axes[1, col].set_xlabel("x") plt.colorbar(im1, ax=axes[1, col], format="%.2e") # φ_E sur grille x (évalué à v=0) x1d = jnp.array(x_lin[:, None]) v0_1d = jnp.zeros_like(x1d) t1d = jnp.full_like(x1d, t_val) vals1d = np.array(pinn.evaluate(x1d, v0_1d, t1d)) phi_E = vals1d[:, 3] axes[2, col].plot(x_lin, phi_E, color="tomato") axes[2, col].axhline(0.0, color="gray", linestyle="--", linewidth=0.8) axes[2, col].set_title(f"φ_E(x, t={t_val:.1f})") axes[2, col].set_xlabel("x") plt.tight_layout() plt.show()