# %% import copy import os import timeit import jax import jax.numpy as jnp import matplotlib.pyplot as plt import numpy as np from scimba_jax.domains.meshless_domains.base import VolumetricDomain from scimba_jax.domains.meshless_domains.domains_2d import Square2D from scimba_jax.nonlinear_approximation.approximation_spaces.densityflow_approximation_spaces import ( DensityFlowApproximationSpace, DensityFlowInvertibleApproximationSpace, ) from scimba_jax.nonlinear_approximation.integration.monte_carlo import ( DomainSampler, TensorizedSampler, ) from scimba_jax.nonlinear_approximation.integration.monte_carlo_time import ( UniformTimeSampler, ) from scimba_jax.nonlinear_approximation.model_class.funcparam_vectorial import ( ParamScalarFunction, ParamVecFunction, ) from scimba_jax.nonlinear_approximation.networks.mlp import MLP from scimba_jax.nonlinear_approximation.networks.structure_preserving_nets.symplectic_nets import ( 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 ( DOMAIN_TYPE, NDARRAYS_FUNC_TYPE, PARAM_FUNC_TYPE, InitialResidual, InteriorResidual, ) # figures directory anchored to this script (works from any CWD) FIG_DIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), "fig") os.makedirs(FIG_DIR, exist_ok=True) class LinearVlasovResidualInvertibleLossFlow(InteriorResidual): def __init__( self, domain: VolumetricDomain, time_domain: tuple[float, float], f_rhs: NDARRAYS_FUNC_TYPE | None = None, Tf: float = 1.0, ): super().__init__( domain=domain, size=2, model_type="x_t", f_rhs=f_rhs, time_domain=time_domain, ) self.Tf = Tf def velocity(self, x): v1 = x[1] v2 = jnp.sin(x[0]) return jnp.array([v1, v2]) def construct_residual( self, phi: PARAM_FUNC_TYPE, inv_phi: PARAM_FUNC_TYPE, rho: PARAM_FUNC_TYPE ) -> PARAM_FUNC_TYPE: phi1, phi2 = inv_phi.components() dx_phi1 = phi1.gradient("x") dx_phi2 = phi2.gradient("x") dt_phi1 = phi1.d_t() dt_phi2 = phi2.d_t() return ParamVecFunction.cat( [dt_phi1 + dx_phi1.dot(self.velocity), dt_phi2 + dx_phi2.dot(self.velocity)] ) class LinearVlasovResidualMLPLossFlow(InteriorResidual): def __init__( self, domain: VolumetricDomain, time_domain: tuple[float, float], f_rhs: NDARRAYS_FUNC_TYPE | None = None, Tf: float = 1.0, ): super().__init__( domain=domain, size=2, model_type="x_t", f_rhs=f_rhs, time_domain=time_domain, ) self.Tf = Tf def velocity(self, x): v1 = x[1] v2 = jnp.sin(x[0]) return jnp.array([v1, v2]) def construct_residual( self, inv_phi: PARAM_FUNC_TYPE, rho: PARAM_FUNC_TYPE ) -> PARAM_FUNC_TYPE: phi1, phi2 = inv_phi.components() dx_phi1 = phi1.gradient("x") dx_phi2 = phi2.gradient("x") dt_phi1 = phi1.d_t() dt_phi2 = phi2.d_t() return ParamVecFunction.cat( [dt_phi1 + dx_phi1.dot(self.velocity), dt_phi2 + dx_phi2.dot(self.velocity)] ) class LinearVlasovResidualInvertibleLossDensity(InteriorResidual): def __init__( self, domain: VolumetricDomain, time_domain: tuple[float, float], f_rhs: NDARRAYS_FUNC_TYPE | None = None, Tf: float = 1.0, ): super().__init__( domain=domain, size=1, model_type="x_t", f_rhs=f_rhs, time_domain=time_domain, ) self.Tf = Tf def velocity(self, x): v1 = x[1] v2 = jnp.sin(x[0]) return jnp.array([v1, v2]) def construct_residual( self, phi: PARAM_FUNC_TYPE, inv_phi: PARAM_FUNC_TYPE, rho: PARAM_FUNC_TYPE ) -> PARAM_FUNC_TYPE: dx_rho = rho.gradient("x") dt_rho = rho.d_t() return ParamVecFunction.cat([dt_rho + dx_rho.dot(self.velocity)]) class LinearVlasovResidualMLPLossDensity(InteriorResidual): def __init__( self, domain: VolumetricDomain, time_domain: tuple[float, float], f_rhs: NDARRAYS_FUNC_TYPE | None = None, Tf: float = 1.0, ): super().__init__( domain=domain, size=1, model_type="x_t", f_rhs=f_rhs, time_domain=time_domain, ) self.Tf = Tf def velocity(self, x): v1 = x[1] v2 = jnp.sin(x[0]) return jnp.array([v1, v2]) def construct_residual( self, inv_phi: PARAM_FUNC_TYPE, rho: PARAM_FUNC_TYPE ) -> PARAM_FUNC_TYPE: dx_rho = rho.gradient("x") dt_rho = rho.d_t() return ParamVecFunction.cat([dt_rho + dx_rho.dot(self.velocity)]) class ProjectionResidualInvertibleLossFlow(InitialResidual): def __init__( self, domain: DOMAIN_TYPE, time_domain: float | tuple[float] = (0.0,), size: int = 2, model_type: str = "x_t", f_rhs: NDARRAYS_FUNC_TYPE | None = None, ): super().__init__( domain=domain, time_domain=( (time_domain,) if isinstance(time_domain, float) else time_domain ), size=size, model_type=model_type, f_rhs=f_rhs, ) def construct_residual( self, phi: PARAM_FUNC_TYPE, inv_phi: PARAM_FUNC_TYPE, rho: PARAM_FUNC_TYPE ) -> PARAM_FUNC_TYPE: assert isinstance(rho, ParamScalarFunction) or isinstance(rho, ParamVecFunction) return inv_phi.set_t_0(self.time_domain[0]) class ProjectionResidualMLPLossFlow(InitialResidual): def __init__( self, domain: DOMAIN_TYPE, time_domain: float | tuple[float] = (0.0,), size: int = 2, model_type: str = "x_t", f_rhs: NDARRAYS_FUNC_TYPE | None = None, ): super().__init__( domain=domain, time_domain=( (time_domain,) if isinstance(time_domain, float) else time_domain ), size=size, model_type=model_type, f_rhs=f_rhs, ) def construct_residual( self, inv_phi: PARAM_FUNC_TYPE, rho: PARAM_FUNC_TYPE ) -> PARAM_FUNC_TYPE: assert isinstance(rho, ParamScalarFunction) or isinstance(rho, ParamVecFunction) return inv_phi.set_t_0(self.time_domain[0]) class ProjectionResidualInvertibleLossDensity(InitialResidual): def __init__( self, domain: DOMAIN_TYPE, time_domain: float | tuple[float] = (0.0,), size: int = 1, model_type: str = "x_t", f_rhs: NDARRAYS_FUNC_TYPE | None = None, ): super().__init__( domain=domain, time_domain=( (time_domain,) if isinstance(time_domain, float) else time_domain ), size=size, model_type=model_type, f_rhs=f_rhs, ) def construct_residual( self, phi: PARAM_FUNC_TYPE, inv_phi: PARAM_FUNC_TYPE, rho: PARAM_FUNC_TYPE ) -> PARAM_FUNC_TYPE: assert isinstance(rho, ParamScalarFunction) or isinstance(rho, ParamVecFunction) return rho.set_t_0(self.time_domain[0]) class ProjectionResidualMLPLossDensity(InitialResidual): def __init__( self, domain: DOMAIN_TYPE, time_domain: float | tuple[float] = (0.0,), size: int = 1, model_type: str = "x_t", f_rhs: NDARRAYS_FUNC_TYPE | None = None, ): super().__init__( domain=domain, time_domain=( (time_domain,) if isinstance(time_domain, float) else time_domain ), size=size, model_type=model_type, f_rhs=f_rhs, ) def construct_residual( self, inv_phi: PARAM_FUNC_TYPE, rho: PARAM_FUNC_TYPE ) -> PARAM_FUNC_TYPE: assert isinstance(rho, ParamScalarFunction) or isinstance(rho, ParamVecFunction) return rho.set_t_0(self.time_domain[0]) class LinearVlasovTimeInterval(AbstractPhysicalModel): def __init__( self, main_domain: VolumetricDomain, time_domain: tuple[float, float], f_rhs: NDARRAYS_FUNC_TYPE | None = None, bc: str = "strong", ic: str = "weak", f_ic_rhs: NDARRAYS_FUNC_TYPE | None = None, type_loss: str = "flow", type_net: str = "mlp", Tf=8.0, ): super().__init__(main_domain=main_domain, time_domain=time_domain) if type_loss == "flow": if type_net == "mlp": residual = LinearVlasovResidualMLPLossFlow ic_residual = ProjectionResidualMLPLossFlow else: residual = LinearVlasovResidualInvertibleLossFlow ic_residual = ProjectionResidualInvertibleLossFlow size = 2 else: if type_net == "mlp": residual = LinearVlasovResidualMLPLossDensity ic_residual = ProjectionResidualMLPLossDensity else: residual = LinearVlasovResidualInvertibleLossDensity ic_residual = ProjectionResidualInvertibleLossDensity size = 1 self.physical_residuals: PHYSICAL_RESIDUALS_TYPE = { self.main_domain.get_label(): residual( domain=main_domain, time_domain=time_domain, f_rhs=f_rhs, Tf=Tf, ), } if ic == "weak": label = "ic " + self.main_domain.get_label() self.physical_residuals[label] = ic_residual( domain=main_domain, time_domain=time_domain, size=size, model_type="x_t", f_rhs=f_ic_rhs, ) def renew_time_domain(self, new_time_domain: tuple[float, float]): """Renew the time domain of the physical model and of all its residuals.""" self.time_domain = new_time_domain for label in self.physical_residuals: self.physical_residuals[label].time_domain = new_time_domain def create_net( type_net, size, conditional_size, key, hidden_size=14, num_layers=8, h=0.5 ): if type_net == "invertible": model = PeriodicGSymplecticNet( size=size, conditional_size=conditional_size, width=hidden_size, nb_layers=num_layers, key=key, period=jnp.array([2.0 * jnp.pi]), activation="tanh", h=h, ) else: model = MLP( in_size=size + conditional_size, hidden_sizes=[hidden_size] * num_layers, out_size=size, activation="tanh", key=key, embedding="periodic", embedding_axes=(0,), periods=(2.0 * jnp.pi,), ) return model def initial_density(x): res = 1.0 / (2.0 * jnp.pi) ** 0.5 * jnp.exp(-0.5 * x[1] ** 2) return res def initial_solution_flow(x): return x def initial_solution_density(x): return initial_density(x) # --- Semi-Lagrangian reference solution (dx/dt = v, dv/dt = sin(x)) --- def sl_characteristic_rhs(state): x, v = state return jnp.array([v, jnp.sin(x)]) def sl_rk4_step(state, dt): k1 = sl_characteristic_rhs(state) k2 = sl_characteristic_rhs(state + 0.5 * dt * k1) k3 = sl_characteristic_rhs(state + 0.5 * dt * k2) k4 = sl_characteristic_rhs(state + dt * k3) return state + (dt / 6.0) * (k1 + 2.0 * k2 + 2.0 * k3 + k4) def sl_backward_flow(x, v, t, n_steps): dt = -t / n_steps def body(state, _): return sl_rk4_step(state, dt), None state, _ = jax.lax.scan(body, jnp.array([x, v]), None, length=n_steps) return state def sl_density(x, v, t, n_steps): _, v0 = sl_backward_flow(x, v, t, n_steps) return initial_density(jnp.array([0.0, v0])) sl_density_grid = jax.jit( jax.vmap( jax.vmap(sl_density, in_axes=(0, 0, None, None)), in_axes=(0, 0, None, None) ), static_argnums=(3,), ) def plot_loss_history(pinn, flow_number): for label in pinn.losses.losses_history: if (label == "total") or (pinn.losses.losses_history[label].shape[1] == 1): plt.semilogy(pinn.losses.losses_history[label], label=label + " loss") else: for i in range(pinn.losses.losses_history[label].shape[1]): plt.semilogy( pinn.losses.losses_history[label][:, i], label=label + "%d loss" % i, ) plt.title(f"loss history for flow {flow_number}") plt.legend() plt.show() def plot_density_pinn_ref_error( name, pinn, type_network=None, n_visu=512, n_steps_sl=400, Tf=4.5 ): id_density = 2 if type_network == "mlp" else 4 sh = (n_visu, n_visu) x1_lin = jnp.linspace(0, 2.0 * jnp.pi, n_visu) x2_lin = jnp.linspace(-6, 6, n_visu) x1, x2 = jnp.meshgrid(x1_lin, x2_lin) x = jnp.stack([x1.ravel(), x2.ravel()], axis=-1) times = [0.0, Tf / 3, 2 * Tf / 3, Tf] rho_pinn_list = [] rho_sl_list = [] for i, t_val in enumerate(times): print(f" Computing PINN + SL at t={t_val:.2f} ({i + 1}/{len(times)}) ...") t = jnp.ones_like(x[:, 0:1]) * t_val rho_pinn = pinn.evaluate(x, t)[:, id_density : id_density + 1].reshape(*sh) rho_sl = sl_density_grid(x1, x2, t_val, n_steps_sl) rho_pinn_list.append(rho_pinn) rho_sl_list.append(rho_sl) print(" Plotting SL reference ...") def plot_density_sl(ax, rho, t_val): im = ax.contourf(x1, x2, rho, levels=256, cmap="turbo", zorder=-20) plt.colorbar(im, ax=ax) ax.contour( im, levels=im.levels[:: 256 // 6], colors="w", alpha=0.5, linewidths=0.8, zorder=-20, ) ax.set_rasterization_zorder(-10) ax.set_title( f"SL density t={t_val} (min={float(rho.min()):.2f}, max={float(rho.max()):.2f})" ) fig, axes = plt.subplots(2, 2, figsize=(10, 8)) for ax, rho, t_val in zip(axes.ravel(), rho_sl_list, times): plot_density_sl(ax, rho, t_val) plt.tight_layout() plt.savefig(os.path.join(FIG_DIR, f"{name}_density_sl.pdf")) plt.show() fig, ax = plt.subplots(1, 1, figsize=(5, 4)) plot_density_sl(ax, rho_sl_list[-1], times[-1]) plt.tight_layout() plt.savefig(os.path.join(FIG_DIR, f"{name}_density_sl_last_time.pdf")) plt.show() print(" Plotting PINN density ...") def plot_density_pinn(ax, rho, t_val): im = ax.contourf(x1, x2, rho, levels=256, cmap="turbo", zorder=-20) plt.colorbar(im, ax=ax) ax.contour( im, levels=im.levels[:: 256 // 6], colors="w", alpha=0.5, linewidths=0.8, zorder=-20, ) ax.set_rasterization_zorder(-10) ax.set_title( f"PINN density t={t_val} (min={float(rho.min()):.2f}, max={float(rho.max()):.2f})" ) fig, axes = plt.subplots(2, 2, figsize=(10, 8)) for ax, rho_pinn, t_val in zip(axes.ravel(), rho_pinn_list, times): plot_density_pinn(ax, rho_pinn, t_val) plt.tight_layout() plt.savefig(os.path.join(FIG_DIR, f"{name}_density_pinn.pdf")) plt.show() fig, ax = plt.subplots(1, 1, figsize=(5, 4)) plot_density_pinn(ax, rho_pinn_list[-1], times[-1]) plt.tight_layout() plt.savefig(os.path.join(FIG_DIR, f"{name}_density_pinn_last_time.pdf")) plt.show() print(" Plotting error |PINN - SL| ...") dx1 = x1_lin[1] - x1_lin[0] dx2 = x2_lin[1] - x2_lin[0] def plot_density_error(ax, rho_pinn, rho_sl, t_val): err = jnp.abs(rho_pinn - rho_sl) l2_err = float(jnp.sqrt(jnp.sum(err**2) * dx1 * dx2)) l2_ref = float(jnp.sqrt(jnp.sum(rho_sl**2) * dx1 * dx2)) rel_l2 = l2_err / l2_ref if l2_ref > 0 else l2_err im = ax.contourf(x1, x2, err, levels=256, cmap="magma", zorder=-20) plt.colorbar(im, ax=ax) ax.contour( im, levels=im.levels[:: 256 // 6], colors="w", alpha=0.5, linewidths=0.8, zorder=-20, ) ax.set_rasterization_zorder(-10) ax.set_title(rf"|PINN - SL|: t={t_val} (rel. $L^2$={float(rel_l2):.2e})") fig, axes = plt.subplots(2, 2, figsize=(10, 8)) for ax, rho_pinn, rho_sl, t_val in zip( axes.ravel(), rho_pinn_list, rho_sl_list, times ): plot_density_error(ax, rho_pinn, rho_sl, t_val) plt.tight_layout() plt.savefig(os.path.join(FIG_DIR, f"{name}_density_error.pdf")) plt.show() fig, ax = plt.subplots(1, 1, figsize=(5, 4)) plot_density_error(ax, rho_pinn_list[-1], rho_sl_list[-1], times[-1]) plt.tight_layout() plt.savefig(os.path.join(FIG_DIR, f"{name}_density_error_last_time.pdf")) plt.show() def plot_last_pinn(name, pinn, nb_models, type_network=None, n_visu=256, Tf=4.5): pinn.space.inference_mode = True pinn.space.idx_current_flow = nb_models - 1 plot_density_pinn_ref_error( name, pinn, type_network=type_network, n_visu=n_visu, Tf=Tf ) pinn.space.inference_mode = False def plot_losses_and_density(name, list_of_pinns, type_network=None, n_visu=256, Tf=4.5): for i, pinn in enumerate(list_of_pinns): plot_loss_history(pinn, flow_number=i) plot_last_pinn( name, list_of_pinns[-1], len(list_of_pinns), type_network, n_visu, Tf ) # %% ========== Benchmark helpers ========== def hamiltonian(x1, x2): return -jnp.cos(x1) + 0.5 * x2**2 def prepare_train( type_net, nb_models, Tf, n_epochs, n_colloc, n_ic_colloc, hidden_size=None, num_layers=None, h=0.05, seed=0, scan_frozen_flows=False, ): dim = 2 params_dim = 0 type_loss = "flow" time_and_params_dim = 1 + params_dim domain_x = [(0.0, 2.0 * jnp.pi), (-6.0, 6.0)] dx = Square2D(domain_x, is_main_domain=True) key = jax.random.PRNGKey(seed) keys = jax.random.split(key, nb_models) if hidden_size is None: hidden_size = 16 if type_net == "invertible" else 12 if num_layers is None: num_layers = 6 if type_net == "invertible" else 4 models = [ create_net( type_net=type_net, size=dim, conditional_size=time_and_params_dim, key=keys[i], hidden_size=hidden_size, num_layers=num_layers, h=h, ) for i in range(nb_models) ] Deltat = Tf / nb_models time_intervals = [(Deltat * i, Deltat * (i + 1)) for i in range(nb_models)] if type_net == "invertible": space = DensityFlowInvertibleApproximationSpace( dim=dim, params_dim=params_dim, models=models, model_type="x_t", initial_density=initial_density, time_intervals=time_intervals, ) else: space = DensityFlowApproximationSpace( dim=dim, params_dim=params_dim, models=models, model_type="x_t", initial_density=initial_density, time_intervals=time_intervals, add_id=True, h=h, scan_frozen_flows=scan_frozen_flows, ) weights = {"interior": [0.3, 0.3], "ic interior": [5.0, 5.0]} initial_domain_t = (0.0, time_intervals[0][1]) sampler = TensorizedSampler( [DomainSampler(dx), UniformTimeSampler(initial_domain_t)], bc=False, ic=True, model_type="x_t", ) model = LinearVlasovTimeInterval( main_domain=dx, time_domain=initial_domain_t, f_ic_rhs=initial_solution_flow, type_loss=type_loss, type_net=type_net, ) return key, time_intervals, weights, model, space, sampler def train_config( type_net, nb_models, Tf, n_epochs, n_colloc, n_ic_colloc, hidden_size=None, num_layers=None, h=0.05, seed=0, scan_frozen_flows=False, ): key, time_intervals, weights, model, space, sampler = prepare_train( type_net=type_net, nb_models=nb_models, Tf=Tf, n_epochs=n_epochs, n_colloc=n_colloc, n_ic_colloc=n_ic_colloc, hidden_size=hidden_size, num_layers=num_layers, h=h, seed=seed, scan_frozen_flows=scan_frozen_flows, ) list_of_pinns = [] total_time = 0.0 for i in range(nb_models): space.idx_current_flow = i domain_t = time_intervals[i] 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=1e-7, gram_solver="cholesky", gram_assembly="gemm", ) start = timeit.default_timer() key, pinn = pinn.project( key, space, n_epochs=n_epochs, n_colloc=n_colloc, n_ic_colloc=n_ic_colloc ) total_time += timeit.default_timer() - start space = pinn.space list_of_pinns.append(copy.deepcopy(pinn)) print( f" Flow {i + 1}/{nb_models}: best loss = {float(pinn.best_loss['total']):.2e}" ) return list_of_pinns, total_time def compute_errors_and_energy( pinn, type_network, Tf, n_visu=1024, n_steps_sl=500, n_times=20, max_points_per_chunk=128**2, ): id_density = 2 if type_network == "mlp" else 4 # Keep endpoint=False exactly as requested for the periodic integration x1_lin, dx1 = jnp.linspace(0, 2.0 * jnp.pi, n_visu, endpoint=False, retstep=True) x2_lin, dx2 = jnp.linspace(-6, 6, n_visu, endpoint=False, retstep=True) # Calculate how to chunk the x2_lin array (rows of the meshgrid) n_rows_per_chunk = max(1, max_points_per_chunk // n_visu) n_chunks = int(np.ceil(n_visu / n_rows_per_chunk)) x2_lin_chunks = jnp.array_split(x2_lin, n_chunks) times = np.linspace(0.0, Tf, n_times) l2_errors, linf_errors, rel_l2_errors = [], [], [] energies_pinn, energies_sl, mass_pinn = [], [], [] for t_val in times: # Accumulators for the spatial integrals and global maximums sum_err2 = 0.0 max_err = 0.0 sum_rho_sl2 = 0.0 sum_e_pinn = 0.0 sum_e_sl = 0.0 sum_m_pinn = 0.0 for x2_lin_c in x2_lin_chunks: # Generate the meshgrid purely for the current chunk x1_c, x2_c = jnp.meshgrid(x1_lin, x2_lin_c) sh_c = x1_c.shape x_c = jnp.stack([x1_c.ravel(), x2_c.ravel()], axis=-1) t_c = jnp.ones_like(x_c[:, 0:1]) * t_val # Evaluate models on this specific chunk rho_pinn_c = pinn.evaluate(x_c, t_c)[ :, id_density : id_density + 1 ].reshape(*sh_c) rho_sl_c = sl_density_grid(x1_c, x2_c, t_val, n_steps_sl) H_c = hamiltonian(x1_c, x2_c) # Sub-grid error metrics err_c = jnp.abs(rho_pinn_c - rho_sl_c) # Accumulate chunk sums and update max bounds sum_err2 += jnp.sum(err_c**2) max_err = jnp.maximum(max_err, err_c.max()) sum_rho_sl2 += jnp.sum(rho_sl_c**2) sum_e_pinn += jnp.sum(H_c * rho_pinn_c) sum_e_sl += jnp.sum(H_c * rho_sl_c) sum_m_pinn += jnp.sum(rho_pinn_c) # Apply the dx1 * dx2 factors to the fully accumulated sums l2_err = float(jnp.sqrt(sum_err2 * dx1 * dx2)) linf_err = float(max_err) l2_ref = float(jnp.sqrt(sum_rho_sl2 * dx1 * dx2)) rel_l2 = l2_err / l2_ref if l2_ref > 0 else l2_err e_pinn = float(sum_e_pinn * dx1 * dx2) e_sl = float(sum_e_sl * dx1 * dx2) m_pinn = float(sum_m_pinn * dx1 * dx2) l2_errors.append(l2_err) linf_errors.append(linf_err) rel_l2_errors.append(rel_l2) energies_pinn.append(e_pinn) energies_sl.append(e_sl) mass_pinn.append(m_pinn) return { "times": times, "l2_errors": np.array(l2_errors), "linf_errors": np.array(linf_errors), "rel_l2_errors": np.array(rel_l2_errors), "energies_pinn": np.array(energies_pinn), "energies_sl": np.array(energies_sl), "mass_pinn": np.array(mass_pinn), } def save_benchmark_plots( pinn, results, label, type_network, Tf, save_dir, n_visu=256, n_steps_sl=500 ): os.makedirs(save_dir, exist_ok=True) times = results["times"] # --- Energy conservation --- fig, ax = plt.subplots(figsize=(8, 4)) ax.plot(times, results["energies_pinn"], "o-", label="PINN", markersize=3) ax.plot(times, results["energies_sl"], "s--", label="SL ref", markersize=3) ax.set_xlabel("t") ax.set_ylabel("Energy ∫ H ρ dx dv") ax.set_title(f"Energy conservation — {label}") ax.legend() fig.tight_layout() fig.savefig(os.path.join(save_dir, f"energy_{label}.pdf")) plt.close(fig) # --- Mass conservation --- fig, ax = plt.subplots(figsize=(8, 4)) ax.plot(times, results["mass_pinn"], "o-", markersize=3) ax.axhline(results["mass_pinn"][0], ls="--", color="gray", label="initial mass") ax.set_xlabel("t") ax.set_ylabel("Mass ∫ ρ dx dv") ax.set_title(f"Mass conservation — {label}") ax.legend() fig.tight_layout() fig.savefig(os.path.join(save_dir, f"mass_{label}.pdf")) plt.close(fig) # --- L² error vs time --- fig, ax = plt.subplots(figsize=(8, 4)) ax.semilogy(times, results["l2_errors"], "o-", markersize=3) ax.set_xlabel("t") ax.set_ylabel("L² error") ax.set_title(f"L² error vs SL reference — {label}") fig.tight_layout() fig.savefig(os.path.join(save_dir, f"l2_error_{label}.pdf")) plt.close(fig) # save everything to a csv data = np.c_[ results["times"], results["l2_errors"], results["linf_errors"], results["rel_l2_errors"], results["energies_pinn"], results["energies_sl"], results["mass_pinn"], ] np.savetxt( os.path.join(save_dir, f"results_{label}.csv"), data, delimiter=",", header="time,l2_error,linf_error,rel_l2_error,energies_pinn,energies_sl,mass_pinn", comments="", ) # --- Density snapshots (PINN vs SL) --- id_density = 2 if type_network == "mlp" else 4 x1_lin = jnp.linspace(0, 2.0 * jnp.pi, n_visu) x2_lin = jnp.linspace(-6, 6, n_visu) x1g, x2g = jnp.meshgrid(x1_lin, x2_lin) x = jnp.stack([x1g.ravel(), x2g.ravel()], axis=-1) sh = (n_visu, n_visu) snap_times = [0.0, Tf / 3, 2 * Tf / 3, Tf] fig, axes = plt.subplots(2, 4, figsize=(20, 8)) for j, t_val in enumerate(snap_times): t = jnp.ones_like(x[:, 0:1]) * t_val rho_pinn = pinn.evaluate(x, t)[:, id_density : id_density + 1].reshape(*sh) rho_sl = sl_density_grid(x1g, x2g, t_val, n_steps_sl) im = axes[0, j].contourf( x1g, x2g, rho_pinn, levels=256, cmap="turbo", zorder=-20 ) plt.colorbar(im, ax=axes[0, j]) axes[0, j].contour( im, levels=im.levels[:: 256 // 6], colors="w", alpha=0.5, linewidths=0.8, zorder=-20, ) axes[0, j].set_title(f"PINN t={t_val:.2f}") axes[0, j].set_rasterization_zorder(-10) error = jnp.abs(rho_pinn - rho_sl) im = axes[1, j].contourf( x1g, x2g, error, levels=256, cmap="inferno", zorder=-20 ) plt.colorbar(im, ax=axes[1, j]) axes[1, j].contour( im, levels=im.levels[:: 256 // 6], colors="w", alpha=0.5, linewidths=0.8, zorder=-20, ) axes[1, j].set_title(f"error at t={t_val:.2f}") axes[1, j].set_rasterization_zorder(-10) fig.suptitle(label, fontsize=14) fig.tight_layout() fig.savefig(os.path.join(save_dir, f"density_{label}.pdf")) plt.close(fig) def run_benchmark( configs=None, Tf=2.0, n_colloc=5_000, n_ic_colloc=3_000, save_dir="benchmark_vlasov" ): """Run benchmark over multiple configurations. Each config in `configs` is a dict with keys: type_net: "mlp" or "invertible" nb_flows: number of time sub-intervals (1, 2, 4, …) n_epochs: training epochs per flow hidden_size: (optional) width of each network num_layers: (optional) depth of each network """ if configs is None: configs = [ {"type_net": "mlp", "nb_flows": 1, "n_epochs": 300}, {"type_net": "mlp", "nb_flows": 2, "n_epochs": 300}, {"type_net": "invertible", "nb_flows": 1, "n_epochs": 300}, {"type_net": "invertible", "nb_flows": 2, "n_epochs": 300}, ] os.makedirs(save_dir, exist_ok=True) all_results = [] for cfg in configs: type_net = cfg["type_net"] nb_flows = cfg["nb_flows"] n_epochs = cfg["n_epochs"] hidden_size = cfg.get("hidden_size") num_layers = cfg.get("num_layers") label = f"{type_net}_{nb_flows}flows_{n_epochs}ep" if hidden_size is not None: label += f"_h{hidden_size}" if num_layers is not None: label += f"_L{num_layers}" print(f"\n{'=' * 60}") print(f" Config: {label}") print(f"{'=' * 60}") pinns, train_time = train_config( type_net, nb_flows, Tf, n_epochs, n_colloc, n_ic_colloc, hidden_size=hidden_size, num_layers=num_layers, ) pinn_last = pinns[-1] pinn_last.space.inference_mode = True pinn_last.space.idx_current_flow = nb_flows - 1 results = compute_errors_and_energy(pinn_last, type_net, Tf) save_benchmark_plots(pinn_last, results, label, type_net, Tf, save_dir) pinn_last.space.inference_mode = False all_results.append( { "config": label, "type_net": type_net, "nb_flows": nb_flows, "train_time": train_time, "best_loss": float(pinns[-1].best_loss["total"]), "l2_err_Tf": results["l2_errors"][-1], "linf_err_Tf": results["linf_errors"][-1], "rel_l2_Tf": results["rel_l2_errors"][-1], "energy_drift": abs( results["energies_pinn"][-1] - results["energies_pinn"][0] ), "mass_drift": abs(results["mass_pinn"][-1] - results["mass_pinn"][0]), "timeseries": results, } ) # --- Summary table --- header = ( f"{'Config':<22s} {'Loss':>10s} {'L²(Tf)':>10s} {'L∞(Tf)':>10s} " f"{'relL²(Tf)':>10s} {'ΔE':>10s} {'ΔM':>10s} {'Time(s)':>10s}" ) print(f"\n{'=' * len(header)}") print(" BENCHMARK SUMMARY") print(f"{'=' * len(header)}") print(header) print("-" * len(header)) for r in all_results: print( f"{r['config']:<22s} {r['best_loss']:10.2e} {r['l2_err_Tf']:10.2e} " f"{r['linf_err_Tf']:10.2e} {r['rel_l2_Tf']:10.2e} " f"{r['energy_drift']:10.2e} {r['mass_drift']:10.2e} {r['train_time']:10.1f}" ) print(f"{'=' * len(header)}") print(f"Plots saved in {save_dir}/") # --- Save summary CSV --- csv_path = os.path.join(save_dir, "benchmark_summary.csv") with open(csv_path, "w") as f: f.write( "config,type_net,nb_flows,best_loss,l2_err_Tf,linf_err_Tf,rel_l2_Tf," "energy_drift,mass_drift,train_time_s\n" ) for r in all_results: f.write( f"{r['config']},{r['type_net']},{r['nb_flows']},{r['best_loss']:.6e}," f"{r['l2_err_Tf']:.6e},{r['linf_err_Tf']:.6e},{r['rel_l2_Tf']:.6e}," f"{r['energy_drift']:.6e},{r['mass_drift']:.6e},{r['train_time']:.1f}\n" ) # --- Save per-time data CSV --- for r in all_results: ts_path = os.path.join(save_dir, f"timeseries_{r['config']}.csv") res = r["timeseries"] with open(ts_path, "w") as f: f.write("t,l2_err,linf_err,rel_l2_err,energy_pinn,energy_sl,mass_pinn\n") for k in range(len(res["times"])): f.write( f"{res['times'][k]:.4f},{res['l2_errors'][k]:.6e}," f"{res['linf_errors'][k]:.6e},{res['rel_l2_errors'][k]:.6e}," f"{res['energies_pinn'][k]:.6e},{res['energies_sl'][k]:.6e}," f"{res['mass_pinn'][k]:.6e}\n" ) print(f"Plots and CSVs saved in {save_dir}/") return all_results # %% ========== Single run ========== name = "fine" fine_sim = True # Set to False to skip training and plotting if fine_sim := True: type_net = "mlp" nb_models = 4 Tf = 4.5 N_EPOCHS = 200 list_of_pinns, train_time = train_config( type_net=type_net, nb_models=nb_models, Tf=Tf, n_epochs=N_EPOCHS, n_colloc=5_000, n_ic_colloc=3_000, hidden_size=12, num_layers=5, h=1.0, scan_frozen_flows=True, ) plot_losses_and_density( name, list_of_pinns, type_network=type_net, n_visu=256, Tf=Tf ) # %% ========== Comparison PINN vs SL at different resolutions ========== def _sl_density_on_grid(n_visu, n_steps, t_val): """Compute SL density on a fresh grid of size n_visu². Returns (rho, x1, x2).""" x1_lin = jnp.linspace(0, 2.0 * jnp.pi, n_visu) x2_lin = jnp.linspace(-6, 6, n_visu) x1, x2 = jnp.meshgrid(x1_lin, x2_lin) return sl_density_grid(x1, x2, t_val, n_steps), x1_lin, x2_lin def compare_pinn_vs_sl( list_of_pinns, type_network, Tf, train_time, n_visu_ref=200, n_steps_sl_ref=500, n_times=10, ): """Compare PINN to 3 SL resolutions (time x space) vs a fine SL reference. Each SL config has its own spatial grid and time-step count. Errors are interpolated onto the reference grid for fair comparison. Computes one time-step at a time to limit RAM usage. """ id_density = 2 if type_network == "mlp" else 4 x1_ref = jnp.linspace(0, 2.0 * jnp.pi, n_visu_ref) x2_ref = jnp.linspace(-6, 6, n_visu_ref) dx1 = float(x1_ref[1] - x1_ref[0]) dx2 = float(x2_ref[1] - x2_ref[0]) x1g, x2g = jnp.meshgrid(x1_ref, x2_ref) x_flat = jnp.stack([x1g.ravel(), x2g.ravel()], axis=-1) sh = (n_visu_ref, n_visu_ref) eval_times = np.linspace(0.0, Tf, n_times + 1)[1:] sl_configs = [ ("SL 200² x 100", 200, 100), ("SL 350² x 200", 350, 200), ("SL 500² x 300", 500, 300), ] # Accumulators: method -> (sum_l2², max_linf) pinn_l2_sum, pinn_linf = 0.0, 0.0 sl_l2_sums = {name: 0.0 for name, _, _ in sl_configs} sl_linfs = {name: 0.0 for name, _, _ in sl_configs} time_pinn_eval = 0.0 time_ref_total = 0.0 sl_times = {name: 0.0 for name, _, _ in sl_configs} pinn_last = list_of_pinns[-1] pinn_last.space.inference_mode = True pinn_last.space.idx_current_flow = len(list_of_pinns) - 1 # --- Warmup: trigger JIT compilation for all methods before timing --- t_warmup = eval_times[0] _ = sl_density_grid(x1g, x2g, t_warmup, n_steps_sl_ref) t_arr_w = jnp.ones_like(x_flat[:, 0:1]) * t_warmup _ = pinn_last.evaluate(x_flat, t_arr_w) for _, n_visu_sl, n_steps_sl in sl_configs: x1_sl = jnp.linspace(0, 2.0 * jnp.pi, n_visu_sl) x2_sl = jnp.linspace(-6, 6, n_visu_sl) x1g_sl, x2g_sl = jnp.meshgrid(x1_sl, x2_sl) _ = sl_density_grid(x1g_sl, x2g_sl, t_warmup, n_steps_sl) jax.block_until_ready(jnp.array(0.0)) print("Warmup done, starting timed evaluation ...") # --- SL reference --- print(f"Computing SL reference ({n_visu_ref}² x {n_steps_sl_ref} steps) ...") rho_ref_list = [] for t_val in eval_times: start = timeit.default_timer() rho_ref = sl_density_grid(x1g, x2g, t_val, n_steps_sl_ref) jax.block_until_ready(rho_ref) time_ref_total += timeit.default_timer() - start rho_ref_list.append(rho_ref) # --- PINN --- print("Evaluating PINN ...") for i, t_val in enumerate(eval_times): start = timeit.default_timer() t_arr = jnp.ones_like(x_flat[:, 0:1]) * t_val rho_pinn = pinn_last.evaluate(x_flat, t_arr)[ :, id_density : id_density + 1 ].reshape(*sh) jax.block_until_ready(rho_pinn) time_pinn_eval += timeit.default_timer() - start err = jnp.abs(rho_pinn - rho_ref_list[i]) pinn_l2_sum += float(jnp.sum(err**2) * dx1 * dx2) pinn_linf = max(pinn_linf, float(err.max())) del rho_pinn, err # --- SL at various resolutions --- for sl_name, n_visu_sl, n_steps_sl in sl_configs: print(f"Computing {sl_name} ...") x1_sl = jnp.linspace(0, 2.0 * jnp.pi, n_visu_sl) x2_sl = jnp.linspace(-6, 6, n_visu_sl) x1g_sl, x2g_sl = jnp.meshgrid(x1_sl, x2_sl) for i, t_val in enumerate(eval_times): start = timeit.default_timer() rho_sl_coarse = sl_density_grid(x1g_sl, x2g_sl, t_val, n_steps_sl) jax.block_until_ready(rho_sl_coarse) sl_times[sl_name] += timeit.default_timer() - start rho_sl_interp = jax.scipy.ndimage.map_coordinates( rho_sl_coarse, [ (x2g - (-6.0)) / (12.0) * (n_visu_sl - 1), (x1g) / (2.0 * jnp.pi) * (n_visu_sl - 1), ], order=1, mode="wrap", ) err = jnp.abs(rho_sl_interp - rho_ref_list[i]) sl_l2_sums[sl_name] += float(jnp.sum(err**2) * dx1 * dx2) sl_linfs[sl_name] = max(sl_linfs[sl_name], float(err.max())) del rho_sl_coarse, rho_sl_interp, err del rho_ref_list pinn_last.space.inference_mode = False # Build rows: (name, l2, linf, solve_time, eval_time) rows = [("PINN", np.sqrt(pinn_l2_sum), pinn_linf, train_time, time_pinn_eval)] for sl_name, _, _ in sl_configs: rows.append( ( sl_name, np.sqrt(sl_l2_sums[sl_name]), sl_linfs[sl_name], sl_times[sl_name], sl_times[sl_name], ) ) # Print table header = ( f"{'Method':<18s} {'L²(sum t)':>12s} {'max L∞':>12s} " f"{'Solve(s)':>10s} {'Eval(s)':>10s}" ) print(f"\n{'=' * len(header)}") print(f" PINN vs SL (ref: SL {n_visu_ref}² × {n_steps_sl_ref} steps)") print(f" {n_times} eval times in (0, {Tf}]") print(f"{'=' * len(header)}") print(header) print("-" * len(header)) for name, l2, linf, ts, te in rows: print(f"{name:<18s} {l2:12.4e} {linf:12.4e} {ts:10.2f} {te:10.2f}") print(f"{'=' * len(header)}") print(f" (SL ref total time: {time_ref_total:.2f}s)") if compare := False: compare_pinn_vs_sl( list_of_pinns, type_network=type_net, Tf=Tf, train_time=train_time ) # %% ========== Benchmark (set bench=True to run) ========== bench = True if bench: benchmark_configs = [ { "type_net": "mlp", "nb_flows": 1, "n_epochs": 600, "hidden_size": 20, "num_layers": 8, }, { "type_net": "mlp", "nb_flows": 2, "n_epochs": 400, "hidden_size": 20, "num_layers": 6, }, { "type_net": "mlp", "nb_flows": 4, "n_epochs": 150, "hidden_size": 20, "num_layers": 4, }, { "type_net": "invertible", "nb_flows": 1, "n_epochs": 600, "hidden_size": 16, "num_layers": 12, }, { "type_net": "invertible", "nb_flows": 2, "n_epochs": 300, "hidden_size": 16, "num_layers": 9, }, { "type_net": "invertible", "nb_flows": 4, "n_epochs": 150, "hidden_size": 16, "num_layers": 6, }, ] benchmark_results = run_benchmark(configs=benchmark_configs, Tf=4.5)