# %% import os import pickle 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, ) # 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) 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 compute_plane_coordinates(plane, n_visu): if plane == "x1-v1": 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 = [ x1.ravel(), jnp.full_like(x1.ravel(), x2_slice), v1.ravel(), jnp.full_like(v1.ravel(), v2_slice), ] elif plane == "x1-x2": x1_lin = jnp.linspace(0, 2.0 * jnp.pi, n_visu) x2_lin = jnp.linspace(0, 2.0 * jnp.pi, n_visu) x1, x2 = jnp.meshgrid(x1_lin, x2_lin) v1_slice = 0.0 v2_slice = 0.0 x = [ x1.ravel(), x2.ravel(), jnp.full_like(x1.ravel(), v1_slice), jnp.full_like(x2.ravel(), v2_slice), ] else: raise ValueError(f"Unknown plane: {plane}") return x, (x1, v1) if plane == "x1-v1" else (x1, x2) def plot_density( pinn, type_network=None, n_visu=512, times=(0.0, 1.5, 3.0, 4.5), plane="x1-v1" ): id_density = dim if type_network == "mlp" else 2 * dim sh = (n_visu, n_visu) x, coords = compute_plane_coordinates(plane, n_visu) x = jnp.stack(x, 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) im = ax.contourf(*coords, rho, levels=256, cmap="turbo", zorder=-20) plt.colorbar(im, ax=ax) ax.contour(*coords, 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 = os.path.join(FIG_DIR, "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 = os.path.join(FIG_DIR, "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), plane="x1-v1", ): id_density = dim if type_network == "mlp" else 2 * dim sh = (n_visu, n_visu) x, coords = compute_plane_coordinates(plane, n_visu) x_ = jnp.stack(x, 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(*x, 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(*coords, rho, levels=256, cmap="turbo", zorder=-20) plt.colorbar(im, ax=ax) ax.contour(*coords, 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 = os.path.join(FIG_DIR, "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 = os.path.join(FIG_DIR, "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(*coords, 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 = os.path.join(FIG_DIR, "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 = os.path.join(FIG_DIR, "density_error.pdf") plt.tight_layout() plt.savefig(filename) plt.show() def plot_losses_and_density( domain_x, pinn, type_network=None, n_visu=256, plane="x1-v1" ): pinn.space.inference_mode = True # plot_density(pinn, type_network=type_network, n_visu=n_visu) plot_density( pinn, type_network=type_network, n_visu=n_visu, times=(4.5,), plane=plane ) # plot_sl_reference_and_error(pinn, type_network=type_network, n_visu=n_visu) plot_sl_reference_and_error( pinn, type_network=type_network, n_visu=n_visu, times=(4.5,), plane=plane ) pinn.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, ) with open("models/model_7.pkl", "rb") as f: space = pickle.load(f) pinn = Projector( model, space, sampler, optimizer="ENG", weights=weights, one_loss_per_residual=True, matrix_regularization=5e-7, ) # %% if plot := True: plot_losses_and_density(dx, pinn, type_network=type_net, n_visu=200, plane="x1-v1") # %%