# %% import copy import os import pickle import timeit from pathlib import Path 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_nd import HypercubeND 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, ) def enable_speedup(cache_dir: Path | None = None): if cache_dir is None: if "JAX_CACHE_DIR" in os.environ: cache_dir = Path(os.environ["JAX_CACHE_DIR"]) else: cache_dir = Path.home() / ".scimba" / "jax_compilation_cache" else: cache_dir = Path(cache_dir).expanduser() jax.config.update("jax_compilation_cache_dir", str(cache_dir)) jax.config.update("jax_persistent_cache_min_entry_size_bytes", -1) # Set the minimum compile time to 0 seconds to cache all compilations, regardless # of their duration (our test functions compile quickly!) jax.config.update("jax_persistent_cache_min_compile_time_secs", 0) # enable_speedup() def linear_vlasov_velocity(x): x1, x2, v1, v2 = x return jnp.array([v1, v2, jnp.sin(x1), jnp.sin(x2)]) 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=4, model_type="x_t", f_rhs=f_rhs, time_domain=time_domain, ) self.Tf = Tf def velocity(self, x): return linear_vlasov_velocity(x) def construct_residual( self, phi: PARAM_FUNC_TYPE, inv_phi: PARAM_FUNC_TYPE, rho: PARAM_FUNC_TYPE ) -> PARAM_FUNC_TYPE: residuals = [] for component in inv_phi.components(): dt_component = component.d_t() dx_component = component.gradient("x") residuals.append(dt_component + dx_component.dot(self.velocity)) return ParamVecFunction.cat(residuals) 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=4, model_type="x_t", f_rhs=f_rhs, time_domain=time_domain, ) self.Tf = Tf def velocity(self, x): return linear_vlasov_velocity(x) def construct_residual( self, inv_phi: PARAM_FUNC_TYPE, rho: PARAM_FUNC_TYPE ) -> PARAM_FUNC_TYPE: residuals = [] for component in inv_phi.components(): dt_component = component.d_t() dx_component = component.gradient("x") residuals.append(dt_component + dx_component.dot(self.velocity)) return ParamVecFunction.cat(residuals) 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): return linear_vlasov_velocity(x) 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): return linear_vlasov_velocity(x) 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 = 4, 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 = 4, 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 = 4 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 ): """Crée un InvertibleNet avec des CouplingLayers.""" print(f"create_net: type_net={type_net!r}") 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, 2.0 * jnp.pi]), activation="tanh", h=h, ) else: model = MLP( in_size=dim + time_and_params_dim, hidden_sizes=[hidden_size] * num_layers, out_size=dim, activation="tanh", key=key, embedding="periodic", embedding_axes=(0, 1), periods=(2.0 * jnp.pi, 2.0 * jnp.pi), ) return model def initial_density(x): x1, x2, v1, v2 = x # spatial_modulation = 1.0 + 0.1 * jnp.cos(x1) * jnp.cos(x2) maxwellian_2v = (1.0 / (2.0 * jnp.pi)) * jnp.exp(-0.5 * (v1**2 + v2**2)) # return spatial_modulation * maxwellian_2v return maxwellian_2v def initial_solution_flow(x): return x def initial_solution_density(x): return initial_density(x) # --- Semi-Lagrangian reference solution (2x2v) --- def sl_characteristic_rhs(state): x1, x2, v1, v2 = state return jnp.array([v1, v2, jnp.sin(x1), jnp.sin(x2)]) 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(x1, x2, v1, v2, t, n_steps): dt = -t / n_steps def body(state, _): return sl_rk4_step(state, dt), None state, _ = jax.lax.scan(body, jnp.array([x1, x2, v1, v2]), None, length=n_steps) return state def sl_density(x1, x2, v1, v2, t, n_steps): state0 = sl_backward_flow(x1, x2, v1, v2, t, n_steps) return initial_density(state0) sl_density_grid = jax.jit( jax.vmap(sl_density, in_axes=(0, 0, 0, 0, None, None)), static_argnums=(5,), ) 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, type_network=None, n_visu=512, times=(0.0, 1.5, 3.0, 4.5)): id_density = dim if type_network == "mlp" else 2 * dim sh = (n_visu, n_visu) x1_lin = jnp.linspace(0, 2.0 * jnp.pi, n_visu) v1_lin = jnp.linspace(-6, 6, n_visu) x1, v1 = jnp.meshgrid(x1_lin, v1_lin) x2_slice = jnp.pi v2_slice = 0.0 x = jnp.stack( [ x1.ravel(), jnp.full_like(x1.ravel(), x2_slice), v1.ravel(), jnp.full_like(v1.ravel(), v2_slice), ], axis=-1, ) def plot_rho(ax, t_val): t = jnp.ones_like(x[:, 0:1]) * t_val rho = pinn.evaluate(x, t)[:, id_density : id_density + 1].reshape(*sh) edge_diff = jnp.abs(rho[:, 0] - rho[:, -1]) rel_diff = edge_diff / (jnp.abs(rho).max() + 1e-12) print( f" periodicity check (t={t_val}): " f"max|rho(x1=0)-rho(x1=2pi)| on slice x2={float(x2_slice):.2f}, v2={float(v2_slice):.2f} " f"= {float(edge_diff.max()):.3e}, " f"relative = {float(rel_diff.max()):.3e}" ) im = ax.contourf(x1, v1, rho, levels=256, cmap="turbo", zorder=-20) plt.colorbar(im, ax=ax) ax.contour(x1, v1, rho, levels=6, colors="white", linewidths=0.5, zorder=-20) ax.set_rasterization_zorder(-10) ax.set_title( f"density t={t_val} (min={float(rho.min()):.2f}, max={float(rho.max()):.2f})" ) ax.set_xlabel("x1") ax.set_ylabel("v1") if times == (4.5,): fig, ax = plt.subplots(1, 1, figsize=(5, 4)) plot_rho(ax, times[0]) filename = "fig/density_pinn_last_time.pdf" else: fig, axes = plt.subplots(2, 2, figsize=(10, 8)) for ax, t_val in zip(axes.ravel(), times): plot_rho(ax, t_val) filename = "fig/density_pinn.pdf" plt.tight_layout() plt.savefig(filename) plt.show() def plot_sl_reference_and_error( pinn, type_network=None, n_visu=384, n_steps_sl=400, times=(0.0, 1.5, 3.0, 4.5) ): id_density = dim if type_network == "mlp" else 2 * dim sh = (n_visu, n_visu) x1_lin = jnp.linspace(0, 2.0 * jnp.pi, n_visu) v1_lin = jnp.linspace(-6, 6, n_visu) x1, v1 = jnp.meshgrid(x1_lin, v1_lin) x2_slice = jnp.pi v2_slice = 0.0 x = jnp.stack( [ x1.ravel(), jnp.full_like(x1.ravel(), x2_slice), v1.ravel(), jnp.full_like(v1.ravel(), v2_slice), ], axis=-1, ) rho_pinn_list = [] rho_sl_list = [] for t_val in 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.ravel(), jnp.full_like(x1.ravel(), x2_slice), v1.ravel(), jnp.full_like(v1.ravel(), v2_slice), t_val, n_steps_sl, ).reshape(*sh) rho_pinn_list.append(rho_pinn) rho_sl_list.append(rho_sl) def plot_density_sl(ax, rho, t_val): im = ax.contourf(x1, v1, rho, levels=256, cmap="turbo", zorder=-20) plt.colorbar(im, ax=ax) ax.contour(x1, v1, rho, levels=6, colors="white", linewidths=0.5, 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})" ) ax.set_xlabel("x1") ax.set_ylabel("v1") if times == (4.5,): fig, ax = plt.subplots(1, 1, figsize=(5, 4)) plot_density_sl(ax, rho_sl_list[0], times[0]) filename = "fig/density_sl_last_time.pdf" else: 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) filename = "fig/density_sl.pdf" plt.tight_layout() plt.savefig(filename) plt.show() def plot_density_error(ax, rho_pinn, rho_sl, t_val): err = jnp.abs(rho_pinn - rho_sl) im = ax.contourf(x1, v1, err, levels=256, cmap="turbo", zorder=-20) ax.set_rasterization_zorder(-10) plt.colorbar(im, ax=ax) ax.set_title(f"|PINN - SL|: t={t_val} (max={float(err.max()):.2e})") ax.set_xlabel("x1") ax.set_ylabel("v1") if times == (4.5,): fig, ax = plt.subplots(1, 1, figsize=(5, 4)) plot_density_error(ax, rho_pinn_list[0], rho_sl_list[0], times[0]) filename = "fig/density_error_last_time.pdf" else: 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) filename = "fig/density_error.pdf" plt.tight_layout() plt.savefig(filename) plt.show() def plot_losses_and_density(domain_x, list_of_pinns, type_network=None, n_visu=256): # Courbes de loss : une par flot. for i, pinn in enumerate(list_of_pinns): plot_loss_history(pinn, flow_number=i) # Densité + erreur SL : sur le DERNIER pinn seulement. Lui seul contient tous # les flots entraînés ; en inference_mode, le flot est choisi dynamiquement # selon t (_get_flow_idx_from_t), donc la composition est correcte sur tout # l'horizon. Évaluer un pinn intermédiaire à un t « futur » composerait des # flots non entraînés → densité fausse. pinn_last = list_of_pinns[-1] pinn_last.space.inference_mode = True pinn_last.space.idx_current_flow = len(list_of_pinns) - 1 plot_density(pinn_last, type_network=type_network, n_visu=n_visu) plot_density(pinn_last, type_network=type_network, n_visu=n_visu, times=(4.5,)) plot_sl_reference_and_error(pinn_last, type_network=type_network, n_visu=n_visu) plot_sl_reference_and_error( pinn_last, type_network=type_network, n_visu=n_visu, times=(4.5,) ) pinn_last.space.inference_mode = False # %% N_COLLOC = 10_000 N_IC_COLLOC = 6_000 N_EPOCHS = 500 domain_x = [ (0.0, 2.0 * jnp.pi), (0.0, 2.0 * jnp.pi), (-6.0, 6.0), (-6.0, 6.0), ] dx = HypercubeND(domain_x, is_main_domain=True) key = jax.random.PRNGKey(0) dim = 4 params_dim = 0 nb_models = 8 keys = jax.random.split(key, nb_models) time_and_params_dim = 1 + params_dim ############# important type_loss = "flow" type_net = "mlp" if type_net == "invertible": hidden_size = 16 num_layers = 8 elif type_net == "mlp": hidden_size = 20 num_layers = 3 # 1084 else: raise ValueError("Incorrect network type") 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=0.05, ) for i in range(nb_models) ] Tf = 4.5 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, ) if type_loss == "flow": weights = {"interior": [0.3, 0.3, 0.3, 0.3], "ic interior": [5.0, 5.0, 5.0, 5.0]} else: weights = {"interior": [0.0], "ic interior": [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, ) list_of_pinns = [] for i in range(nb_models): space.idx_current_flow = i domain_t = time_intervals[i] print( f"Training flow {i + 1}/{nb_models} on time interval {domain_t} with {type_loss} loss and {type_net} network" ) 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=5e-7, ) start = timeit.default_timer() key, pinn = pinn.project( key, space, n_epochs=N_EPOCHS, n_colloc=N_COLLOC, n_ic_colloc=N_IC_COLLOC ) end = timeit.default_timer() # 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)) print("best loss: ", pinn.best_loss) print("time for %d epochs: " % N_EPOCHS, end - start) with open("models/model.pkl", "wb") as f: pickle.dump(space, f) # %% if plot := True: plot_losses_and_density(dx, list_of_pinns, type_network=type_net, n_visu=200) # %%