# %% import copy import timeit import jax import jax.numpy as jnp import matplotlib.pyplot as plt from matplotlib.lines import Line2D 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 ( # noqa: E501 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.affine_ode_layers import ( # noqa: E501 AffineFlowLayer, ) from scimba_jax.nonlinear_approximation.networks.structure_preserving_nets.coupling_layers import ( # noqa: E501 CouplingLayer, ) from scimba_jax.nonlinear_approximation.networks.structure_preserving_nets.invertible_nn import ( # noqa: E501 InvertibleNet, ) 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, ) class RotatingTransportResidualInvertibleLossFlow(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, t): A = self.Tf v1 = ( -jnp.cos(jnp.pi * t[0] / A) * jnp.sin(jnp.pi * x[0]) ** 2.0 * jnp.sin(2.0 * jnp.pi * x[1]) ) v2 = ( jnp.cos(jnp.pi * t[0] / A) * jnp.sin(jnp.pi * x[1]) ** 2.0 * jnp.sin(2.0 * jnp.pi * 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 RotatingTransportResidualMLPLossFlow(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, t): A = self.Tf v1 = ( -jnp.cos(jnp.pi * t[0] / A) * jnp.sin(jnp.pi * x[0]) ** 2.0 * jnp.sin(2.0 * jnp.pi * x[1]) ) v2 = ( jnp.cos(jnp.pi * t[0] / A) * jnp.sin(jnp.pi * x[1]) ** 2.0 * jnp.sin(2.0 * jnp.pi * 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 RotatingTransportResidualInvertibleLossDensity(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, t): A = self.Tf v1 = ( -jnp.cos(jnp.pi * t[0] / A) * jnp.sin(jnp.pi * x[0]) ** 2.0 * jnp.sin(2.0 * jnp.pi * x[1]) ) v2 = ( jnp.cos(jnp.pi * t[0] / A) * jnp.sin(jnp.pi * x[1]) ** 2.0 * jnp.sin(2.0 * jnp.pi * 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 RotatingTransportResidualMLPLossDensity(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, t): A = self.Tf v1 = ( -jnp.cos(jnp.pi * t[0] / A) * jnp.sin(jnp.pi * x[0]) ** 2.0 * jnp.sin(2.0 * jnp.pi * x[1]) ) v2 = ( jnp.cos(jnp.pi * t[0] / A) * jnp.sin(jnp.pi * x[1]) ** 2.0 * jnp.sin(2.0 * jnp.pi * 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 RotatingTransportTimeInterval(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 = RotatingTransportResidualMLPLossFlow ic_residual = ProjectionResidualMLPLossFlow else: residual = RotatingTransportResidualInvertibleLossFlow ic_residual = ProjectionResidualInvertibleLossFlow size = 2 else: if type_net == "MLP": residual = RotatingTransportResidualMLPLossDensity ic_residual = ProjectionResidualMLPLossDensity else: residual = RotatingTransportResidualInvertibleLossDensity 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): """Crée un InvertibleNet avec des CouplingLayers.""" if type_net == "invertible": layers_list = [ CouplingLayer( size=size, conditional_size=conditional_size, num_splits=2, ode_layer_type=AffineFlowLayer, hidden_sizes=[hidden_size], activation="tanh", key=key, ) for _ in range(num_layers) ] model = InvertibleNet( size=size, conditional_size=conditional_size, layers_list=layers_list, ) else: model = MLP( in_size=dim + time_and_params_dim, hidden_sizes=[hidden_size] * num_layers, out_size=dim, activation="sine", key=key, final_bias=False, ) return model def initial_density(x): r2 = (x[0:1] - 0.5) ** 2 + (x[1:2] - 0.75) ** 2 return r2 - 0.15**2 def initial_solution_flow(x): return x def initial_solution_density(x): r2 = (x[0:1] - 0.5) ** 2 + (x[1:2] - 0.75) ** 2 return r2 - 0.15**2 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(domain_x, pinn, type_network, n_visu=256, Tf=8.0): if type_network == "MLP": id_density = 2 else: id_density = 4 sh = (n_visu, n_visu) x1_lin = jnp.linspace(domain_x[0][0], domain_x[0][1], n_visu) x2_lin = jnp.linspace(domain_x[1][0], domain_x[1][1], n_visu) x1, x2 = jnp.meshgrid(x1_lin, x2_lin) x = jnp.stack([x1.ravel(), x2.ravel()], axis=-1) xb_fn = jax.vmap( pinn.space.compute_flow_backward_compose_continuous_inference, in_axes=(None, 0, 0), ) det_fn = jax.vmap( pinn.space.abs_det_jacobian_compose_continuous_inference, in_axes=(None, 0, 0), ) def plot_rho(a, title, t_val): t = jnp.ones_like(x[:, 0:1]) * t_val rho = pinn.evaluate(x, t)[:, id_density : id_density + 1].reshape(*sh) x_backward = xb_fn(pinn.space, x, t) if type_network == "invertible": det_jac = det_fn(pinn.space, x_backward, t) else: det_jac = det_fn(pinn.space, x, t) print( f"{title}: x_backward range = " f"[{float(x_backward[:, 0].min()):.2f}, {float(x_backward[:, 0].max()):.2f}] x " f"[{float(x_backward[:, 1].min()):.2f}, {float(x_backward[:, 1].max()):.2f}], " f"det_jac range = [{float(det_jac.min()):.3f}, {float(det_jac.max()):.3f}]" ) im = a.contourf(x1, x2, rho, levels=64, cmap="turbo") plt.colorbar(im, ax=a) a.set_aspect("equal") a.contour(x1, x2, rho, levels=[0.0], colors="white") a.set_aspect("equal") min_rho = float(rho.min()) max_rho = float(rho.max()) title = f"{title} (min={min_rho:.2f}, max={max_rho:.2f})" a.set_title(title) fig, ax = plt.subplots(3, 2, figsize=(12, 18)) t_vals = [float(t_val) for t_val in jnp.linspace(0.0, Tf, 6)] for a, t_val in zip(ax.flat, t_vals): plot_rho(a, f"t={t_val:.2f}", t_val) plt.tight_layout() plt.show() def plot_density_flow_alone(domain_x, pinn, flow_idx, t_vals, n_visu=256): """Plot rho for a single flow evaluated alone (no composition with other flows). Uses model_{flow_idx} directly: x_backward = T_idx^{-1}(x,t), det_jac = |det J_{T_idx^{-1}}(x,t)|, rho = initial_density(x_backward) * det_jac. """ space = pinn.space model = space.models[flow_idx] sh = (n_visu, n_visu) x1_lin = jnp.linspace(domain_x[0][0], domain_x[0][1], n_visu) x2_lin = jnp.linspace(domain_x[1][0], domain_x[1][1], n_visu) x1, x2 = jnp.meshgrid(x1_lin, x2_lin) x = jnp.stack([x1.ravel(), x2.ravel()], axis=-1) mu = jnp.zeros((x.shape[0], 0)) apply_fn = jax.vmap(space._apply_model, in_axes=(None, 0, 0)) det_fn = jax.vmap(space._abs_det_jac_wrt_x, in_axes=(None, 0, 0)) def plot_rho(a, title, t_val): t = jnp.ones_like(x[:, 0:1]) * t_val extra = jnp.concatenate([t, mu], axis=-1) x_backward = apply_fn(model, x, extra) det_jac = det_fn(model, x, extra) rho = (jax.vmap(initial_density)(x_backward)[:, 0] * det_jac).reshape(*sh) print( f"{title} (flow{flow_idx} alone): x_backward range = " f"[{float(x_backward[:, 0].min()):.2f}, {float(x_backward[:, 0].max()):.2f}] x " f"[{float(x_backward[:, 1].min()):.2f}, {float(x_backward[:, 1].max()):.2f}], " f"det_jac range = [{float(det_jac.min()):.3f}, {float(det_jac.max()):.3f}]" ) im = a.contourf(x1, x2, rho, levels=64, cmap="turbo") plt.colorbar(im, ax=a) a.set_aspect("equal") a.contour(x1, x2, rho, levels=[0.0], colors="white") min_rho, max_rho = float(rho.min()), float(rho.max()) a.set_title( f"{title} flow{flow_idx} alone (min={min_rho:.2f}, max={max_rho:.2f})" ) fig, ax = plt.subplots(1, len(t_vals), figsize=(6 * len(t_vals), 6)) if len(t_vals) == 1: ax = [ax] for a, t_val in zip(ax, t_vals): plot_rho(a, f"t={t_val}", t_val) plt.tight_layout() plt.show() def plot_level_set_comparison(domain_x, pinn, type_network, Tf, n_visu=256, n_times=5): if type_network == "MLP": id_density = 2 else: id_density = 4 x1_lin = jnp.linspace(domain_x[0][0], domain_x[0][1], n_visu) x2_lin = jnp.linspace(domain_x[1][0], domain_x[1][1], n_visu) x1, x2 = jnp.meshgrid(x1_lin, x2_lin) x = jnp.stack([x1.ravel(), x2.ravel()], axis=-1) sh = (n_visu, n_visu) def rho_at(t_val): t = jnp.ones_like(x[:, 0:1]) * t_val return pinn.evaluate(x, t)[:, id_density : id_density + 1].reshape(*sh) t_vals = [float(t_val) for t_val in jnp.linspace(0.0, Tf, n_times)] cmap = plt.get_cmap("viridis") fig, ax = plt.subplots(figsize=(6, 6)) for i, t_val in enumerate(t_vals): rho = rho_at(t_val) ax.contour( x1, x2, rho, levels=[0.0], colors=[cmap(i / (n_times - 1))], linewidths=2 ) ax.set_aspect("equal") ax.set_title(f"Level set rho=0 for t in [0, {Tf}]") sm = plt.cm.ScalarMappable(cmap=cmap, norm=plt.Normalize(vmin=0.0, vmax=Tf)) sm.set_array([]) plt.colorbar(sm, ax=ax, label="t") plt.tight_layout() plt.show() def plot_level_set_multi_runs( domain_x, runs, type_network, Tf, n_visu=256, zoom_pad=0.05 ): """Compare the rho=0 level set of several runs at t = Tf/2, 3Tf/4, Tf. ``runs`` is a list of dicts with keys "pinn", "marker", "label" and the optional flag "line" (draw a solid contour line without markers instead of scatter markers). The same color encodes a given time across runs, the marker encodes the run. A zoomed inset of the t=Tf level sets is added in the corner, with padding controlled by ``zoom_pad`` (fraction of the bounding box size). """ if type_network == "MLP": id_density = 2 else: id_density = 4 x1_lin = jnp.linspace(domain_x[0][0], domain_x[0][1], n_visu) x2_lin = jnp.linspace(domain_x[1][0], domain_x[1][1], n_visu) x1, x2 = jnp.meshgrid(x1_lin, x2_lin) x = jnp.stack([x1.ravel(), x2.ravel()], axis=-1) sh = (n_visu, n_visu) t_vals = [Tf / 2.0, 3.0 * Tf / 4.0, Tf] cmap = plt.get_cmap("viridis") colors = [cmap(k / (len(t_vals) - 1)) for k in range(len(t_vals))] fig, ax = plt.subplots(figsize=(6, 6)) def draw_curves(target_ax, t_val, color): t = jnp.ones_like(x[:, 0:1]) * t_val segs_all = [] for run in runs: pinn = run["pinn"] rho = pinn.evaluate(x, t)[:, id_density : id_density + 1].reshape(*sh) is_line = run.get("line", False) linewidths = 1.5 if is_line else 0 cs = target_ax.contour( x1, x2, rho, levels=[0.0], colors=[color], linewidths=linewidths ) for seg in cs.allsegs[0]: if not is_line: target_ax.scatter( seg[::4, 0], seg[::4, 1], marker=run["marker"], color=[color], s=12, ) segs_all.append(seg) return segs_all last_segs = [] for t_val, color in zip(t_vals, colors): segs = draw_curves(ax, t_val, color) if t_val == Tf: last_segs = segs ax.set_aspect("equal") ax.set_title("Level set rho=0 at t=Tf/2, 3Tf/4, Tf") # skip the zoom inset when the rho=0 level set is empty (under-trained run) if last_segs and sum(seg.shape[0] for seg in last_segs) > 0: pts = jnp.concatenate(last_segs, axis=0) xmin, xmax = float(pts[:, 0].min()), float(pts[:, 0].max()) ymin, ymax = float(pts[:, 1].min()), float(pts[:, 1].max()) pad_x = zoom_pad * max(xmax - xmin, 1e-3) pad_y = zoom_pad * max(ymax - ymin, 1e-3) axins = ax.inset_axes([0.55, 0.55, 0.42, 0.42]) draw_curves(axins, Tf, colors[-1]) axins.set_xlim(xmin - pad_x, xmax + pad_x) axins.set_ylim(ymin - pad_y, ymax + pad_y) axins.set_aspect("equal") axins.set_xticks([]) axins.set_yticks([]) ax.indicate_inset_zoom(axins, edgecolor="black") time_legend = [ Line2D( [0], [0], marker="o", color="w", markerfacecolor=colors[k], markersize=8, label=f"t={t_vals[k]:.2f}", ) for k in range(len(t_vals)) ] run_legend = [ Line2D( [0], [0], marker=run["marker"], color="black", linestyle="None", label=run["label"], ) for run in runs ] legend1 = ax.legend(handles=time_legend, loc="upper left", title="time") ax.add_artist(legend1) ax.legend(handles=run_legend, loc="upper right", title="run") plt.tight_layout() plt.show() def print_intermediate_range(pinn, domain_x, t_val, n_visu=64): """Print the range of the last flow's output at t=t_val. This is the value fed as input ``x`` to the previous flow during backward composition: if it falls outside ``domain_x``, the previous flow is evaluated out-of-distribution. """ space = pinn.space x1_lin = jnp.linspace(domain_x[0][0], domain_x[0][1], n_visu) x2_lin = jnp.linspace(domain_x[1][0], domain_x[1][1], n_visu) x1, x2 = jnp.meshgrid(x1_lin, x2_lin) x = jnp.stack([x1.ravel(), x2.ravel()], axis=-1) t = jnp.ones((x.shape[0], 1)) * t_val def apply_last_flow(x_, t_): return space._apply_model(space.models[-1], x_, t_) out = jax.vmap(apply_last_flow)(x, t) print( f"t={t_val}: input x range: x1=[{x[:, 0].min():.3f}, {x[:, 0].max():.3f}], " f"x2=[{x[:, 1].min():.3f}, {x[:, 1].max():.3f}]" ) print( f"t={t_val}: output (input to previous flow) range: " f"x1=[{out[:, 0].min():.3f}, {out[:, 0].max():.3f}], " f"x2=[{out[:, 1].min():.3f}, {out[:, 1].max():.3f}]" ) def plot_losses_and_density( domain_x, list_of_pinns, type_network="MLP", n_visu=256, Tf=8.0 ): for i, pinn in enumerate(list_of_pinns): plot_loss_history(pinn, flow_number=i) pinn = list_of_pinns[-1] pinn.space.inference_mode = True plot_density(domain_x, pinn, type_network=type_network, n_visu=n_visu, Tf=Tf) if type_network == "MLP" and len(pinn.space.models) > 1: for idx in range(len(pinn.space.models)): t0, t1 = pinn.space.time_intervals[idx] t_vals = [float(t0 + frac * (t1 - t0)) for frac in (0.125, 0.5, 1.0)] plot_density_flow_alone( domain_x, pinn, flow_idx=idx, t_vals=t_vals, n_visu=n_visu ) plot_level_set_comparison( domain_x, pinn, type_network=type_network, Tf=Tf, n_visu=n_visu ) # pinn.space.inference_mode = False def plot_single_flow( domain_x, pinn, flow_number, t_vals, type_network="MLP", n_visu=256 ): sh = (n_visu, n_visu) x1_lin = jnp.linspace(domain_x[0][0], domain_x[0][1], n_visu) x2_lin = jnp.linspace(domain_x[1][0], domain_x[1][1], n_visu) x1, x2 = jnp.meshgrid(x1_lin, x2_lin) x = jnp.stack([x1.ravel(), x2.ravel()], axis=-1) if type_network == "MLP": id_flow_backward = 0 else: id_flow_backward = 2 def plot_flow(ax, title, t_val): t = jnp.ones_like(x[:, 0:1]) * t_val ev = pinn.evaluate(x, t) flow_bwd = ev[:, id_flow_backward : id_flow_backward + 2] norm_flow_bwd = jnp.linalg.norm(flow_bwd, axis=-1) im = ax[0].contourf(x1, x2, norm_flow_bwd.reshape(*sh), levels=64, cmap="turbo") plt.colorbar(im, ax=ax[0]) ax[0].set_aspect("equal") ax[0].set_title(f"flow {flow_number}, {title}: (norm)") im = ax[1].contourf( x1, x2, flow_bwd[:, 0].reshape(*sh), levels=64, cmap="turbo" ) plt.colorbar(im, ax=ax[1]) ax[1].set_aspect("equal") ax[1].set_title(f"flow {flow_number}, {title}: (x)") im = ax[2].contourf( x1, x2, flow_bwd[:, 1].reshape(*sh), levels=64, cmap="turbo" ) plt.colorbar(im, ax=ax[2]) ax[2].set_aspect("equal") ax[2].set_title(f"flow {flow_number}, {title}: (y)") fig, ax = plt.subplots(3, 3, figsize=(9, 6)) for a, t_val in zip(ax, t_vals): plot_flow(a, f"t={t_val}", t_val) plt.tight_layout() plt.show() def plot_flows(domain_x, pinn, type_network="MLP", n_visu=256): for i, pinn in enumerate(pinn): t0, t1 = pinn.space.time_intervals[i] t_vals = [float(t0), float((t0 + t1) / 2), float(t1)] plot_single_flow( domain_x, pinn, flow_number=i, t_vals=t_vals, type_network=type_network, n_visu=n_visu, ) # %% N_COLLOC = 6_000 N_IC_COLLOC = 3_000 domain_x = [(0.0, 1.0), (0.0, 1.0)] dx = Square2D(domain_x, is_main_domain=True) key = jax.random.PRNGKey(12) dim = 2 params_dim = 0 time_and_params_dim = 1 + params_dim ############# important type_loss = "flow" type_net = "MLP" Tf = 8.0 def run_training(time_intervals, n_epochs, hidden_size, num_layers, key): nb_models = len(time_intervals) keys = jax.random.split(key, nb_models) 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, ) 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=False, ) if type_loss == "flow": weights = {"interior": [0.3, 0.3], "ic interior": [10.0, 10.0]} else: weights = {"interior": [0.3], "ic interior": [10.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 = RotatingTransportTimeInterval( main_domain=dx, time_domain=initial_domain_t, f_ic_rhs=initial_solution_flow, Tf=8.0, 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, weights=weights, one_loss_per_residual=True, matrix_regularization=1e-6 * (i + 1), gram_solver="cholesky", linesearch="strong-wolfe", # "logarithmic_grid" 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 ) end = timeit.default_timer() list_of_pinns.append(copy.deepcopy(pinn)) space = pinn.space print("best loss: ", pinn.best_loss) print("time for %d epochs: " % n_epochs, end - start) return list_of_pinns, key list_of_pinns_1flow, key = run_training( time_intervals=[(0.0, Tf)], n_epochs=800, hidden_size=16, num_layers=4, key=key ) list_of_pinns_2flow, key = run_training( time_intervals=[(0.0, 4.0), (4.0, Tf)], n_epochs=400, hidden_size=20, num_layers=4, key=key, ) # %% pinn_1flow = list_of_pinns_1flow[-1] pinn_1flow.space.inference_mode = True pinn_2flow = list_of_pinns_2flow[-1] pinn_2flow.space.inference_mode = True print_intermediate_range(pinn_2flow, domain_x, t_val=Tf) print_intermediate_range(pinn_2flow, domain_x, t_val=3.0 * Tf / 4.0) # plot_losses_and_density(domain_x, list_of_pinns_1flow, type_network=type_net, n_visu=256, Tf=Tf) # plot_flows(domain_x, list_of_pinns_1flow, type_network=type_net, n_visu=256) # plot_losses_and_density(domain_x, list_of_pinns_2flow, type_network=type_net, n_visu=256, Tf=Tf) # plot_flows(domain_x, list_of_pinns_2flow, type_network=type_net, n_visu=256) plot_level_set_multi_runs( domain_x, runs=[ {"pinn": pinn_1flow, "marker": "x", "label": "1 flow", "line": True}, {"pinn": pinn_2flow, "marker": "^", "label": "2 flows"}, ], type_network=type_net, Tf=Tf, n_visu=256, ) # %%