r"""Density-to-density transport — strong BC via boundary post-processing. Same problem as density_transport_2d.py but with hard boundary conditions: the post-processing T_post(x) = T(x) - Σ_i e^{-φ_i/ε}(n_i⊗n_i)(T(x)-x) is applied inside the Monge-Ampère residual itself. No boundary loss is needed: T_post ≈ x on ∂Ω by construction. The PDE becomes: g(T_post(x)) det(J_{T_post}(x)) = f(x) where J_{T_post} is the Jacobian of T_post (not ∇²u). """ import timeit import jax import jax.numpy as jnp import matplotlib.pyplot as plt import numpy as np from scimba_jax.domains.meshless_domains.domains_2d import Square2D from scimba_jax.nonlinear_approximation.approximation_spaces.approximation_spaces import ( ApproximationSpace, ) from scimba_jax.nonlinear_approximation.integration.monte_carlo import ( DomainSampler, TensorizedSampler, ) from scimba_jax.nonlinear_approximation.integration.monte_carlo_parameters import ( UniformParametricSampler, ) from scimba_jax.nonlinear_approximation.networks.icnn import ICNN from scimba_jax.nonlinear_approximation.networks.mlp import MLP from scimba_jax.nonlinear_approximation.networks.structure_preserving_nets.invertiblenet_based_icnn import ( InvertibleNetBasedICNN, ) 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 ( PARAM_FUNC_TYPE, InteriorResidual, ) from scimba_jax.physical_models.elliptic_pde.monge_ampere import ( InverseFlowMongeAmpere2D, boundary_post_processing, ) from scimba_jax.physical_models.function_approximator.function_approximator import ( FunctionApproximator, ) # ── Configuration ───────────────────────────────────────────────────────────── N_COLLOC = 6_000 N_EPOCHS = 800 N_INV_EPOCHS = 500 N_NEWTON = 5 EPS_POST = 0.012 # ── Level sets for [0,1]² ───────────────────────────────────────────────────── level_sets_square = [ lambda x: x[0], lambda x: 1.0 - x[0], lambda x: x[1], lambda x: 1.0 - x[1], ] # ── Densities ───────────────────────────────────────────────────────────────── def f(x): return jnp.array([1.0]) def _g_unnorm(y): return jnp.array( [ 1.0 + 5.0 * jnp.exp(-100.0 * jnp.abs((y[0] - 0.5) ** 2 + (y[1] - 0.5) ** 2 - 0.09)) ] ) _key_zg = jax.random.PRNGKey(42) _pts_zg = jax.random.uniform(_key_zg, (200_000, 2)) Z_g = float(jnp.mean(jax.vmap(lambda y: _g_unnorm(y))(_pts_zg))) def g(y): return _g_unnorm(y) / Z_g # ── Monge-Ampère residual with built-in post-processing ────────────────────── class MongeAmperePostProcessedResidual(InteriorResidual): """g(T_post(x)) det(J_{T_post}(x)) - f(x) = 0 where T_post(x) = boundary_post_processing(∇u(x), x, level_sets). """ f: object g: object level_sets: object eps_post: float = 0.02 def __init__(self, domain, f, g, level_sets, eps_post=0.02, model_type="x_mu"): super().__init__(domain=domain, size=1, model_type=model_type) self.f = f self.g = g self.level_sets = level_sets self.eps_post = eps_post def construct_residual(self, *vars: PARAM_FUNC_TYPE) -> PARAM_FUNC_TYPE: u = vars[0] grad_u = u.gradient("x") _level_sets = self.level_sets _eps = self.eps_post from scimba_jax.nonlinear_approximation.model_class.funcparam_vectorial import ( ParamFieldFunction, ParamScalarFunction, ) t_post = ParamFieldFunction( u.dims, lambda *args: boundary_post_processing( grad_u(*args), args[1], _level_sets, _eps ), u.f_type, "x", ) g_of_t = self.g << t_post jac_t_post = t_post.jacobian("x") def det_jac_fn(*args): J = jac_t_post(*args) return jnp.array([jnp.linalg.det(J)]) det_jac = ParamScalarFunction(u.dims, det_jac_fn, u.f_type) return g_of_t * det_jac - self.f class MongeAmpere2DStrongBC(AbstractPhysicalModel): """Monge-Ampère with hard BC via boundary post-processing. No boundary loss — T_post enforces T(x) ≈ x on ∂Ω. """ def __init__(self, main_domain, f, g, level_sets, eps_post=0.02, model_type="x_mu"): super().__init__(main_domain=main_domain) self.physical_residuals: PHYSICAL_RESIDUALS_TYPE = { self.main_domain.get_label(): MongeAmperePostProcessedResidual( domain=main_domain, f=f, g=g, level_sets=level_sets, eps_post=eps_post, model_type=model_type, ), } # ── Domain ──────────────────────────────────────────────────────────────────── key = jax.random.PRNGKey(0) domain_mu: list = [] dx = Square2D([(0.0, 1.0), (0.0, 1.0)], is_main_domain=True) dx2 = Square2D([(-2.0, 2.0), (-2.0, 2.0)], is_main_domain=True) # ── ICNN + pré-entraînement ────────────────────────────────────────────────── key, subkey = jax.random.split(key) nn_u = ICNN( in_size=2, out_size=1, hidden_sizes=[16] * 3, activation="softplus", key=subkey ) space_u = ApproximationSpace({"x": 2}, [(nn_u, "scalar", 1)], model_type="x_mu") print("Pré-entraînement u ≈ ½|x|² ...") key, _proj_u = Projector( FunctionApproximator( main_domain=dx2, size=1, model_type="x_mu", f_rhs=lambda *args: jnp.array([0.5 * jnp.dot(args[0], args[0])]), ), space_u, TensorizedSampler( [DomainSampler(dx2), UniformParametricSampler(domain_mu)], bc=False ), ).project(key, space_u, 5, N_COLLOC) space_u = _proj_u.space print("Pré-entraînement terminé.") # ══════════════════════════════════════════════════════════════════════════════ # PHASE 1 — Monge-Ampère with strong BC (no boundary loss) # ══════════════════════════════════════════════════════════════════════════════ print("\n── Phase 1 : Monge-Ampère (strong BC) ──────────────────────────────") model_ma = MongeAmpere2DStrongBC( dx, f=f, g=g, level_sets=level_sets_square, eps_post=EPS_POST ) sampler_ma = TensorizedSampler( [DomainSampler(dx), UniformParametricSampler(domain_mu)], bc=False ) pinn_ma = Projector( model_ma, space_u, sampler_ma, weights={"interior": [1.0]}, matrix_regularization=2.0e-6, ) key, sample_dict = sampler_ma.sample(key, N_COLLOC) print(f" initial loss: {pinn_ma.evaluate_loss(space_u, sample_dict):.4e}") t0 = timeit.default_timer() key, pinn_ma = pinn_ma.project(key, space_u, N_EPOCHS, N_COLLOC) print( f" best loss: {pinn_ma.best_loss['total']:.4e}" f" | {timeit.default_timer() - t0:.1f}s" ) # ══════════════════════════════════════════════════════════════════════════════ # PHASE 2 — Inverse map v ≈ T_post⁻¹ # ══════════════════════════════════════════════════════════════════════════════ print("\n── Phase 2 : inverse map v(T_post(x)) = x ─────────────────────────") key, subkey = jax.random.split(key) nn_v = MLP(in_size=2, out_size=2, hidden_sizes=[16, 16, 16], key=subkey) space_v = ApproximationSpace({"x": 2}, [(nn_v, "field", 2)], model_type="x_mu") model_inv = InverseFlowMongeAmpere2D( dx, space_u=pinn_ma.space, level_sets=level_sets_square, eps_post=EPS_POST ) sampler_inv = TensorizedSampler( [DomainSampler(dx), UniformParametricSampler(domain_mu)], bc=False ) pinn_inv_proj = Projector(model_inv, space_v, sampler_inv) t0 = timeit.default_timer() key, pinn_inv = pinn_inv_proj.project(key, space_v, N_INV_EPOCHS, N_COLLOC) print( f" best loss: {pinn_inv.best_loss['total']:.4e}" f" | {timeit.default_timer() - t0:.1f}s" ) # ══════════════════════════════════════════════════════════════════════════════ # Assembly # ══════════════════════════════════════════════════════════════════════════════ icnn_map = InvertibleNetBasedICNN( u=pinn_ma.space.models[0], size=2, v=pinn_inv.space.models[0], n_newton=N_NEWTON, level_sets=level_sets_square, eps_post=EPS_POST, ) T_forward = jax.jit(jax.vmap(icnn_map)) T_backward = jax.jit(jax.vmap(icnn_map.backward)) # ── Plots ───────────────────────────────────────────────────────────────────── N_GRID = 100 xl = np.linspace(0.0, 1.0, N_GRID) X0, X1 = np.meshgrid(xl, xl) xy_flat = jnp.array(np.stack([X0.ravel(), X1.ravel()], axis=1)) G_grid = np.vectorize(lambda y0, y1: float(g(jnp.array([y0, y1]))[0]))(X0, X1) T_vals = jax.device_get(T_forward(xy_flat)) TX = T_vals[:, 0].reshape(N_GRID, N_GRID) TY = T_vals[:, 1].reshape(N_GRID, N_GRID) _n_lines = 40 _n_pts = 200 def _plot_mesh_fn(ax, fn, color_h="b", color_v="r", title=""): _t = np.linspace(0.0, 1.0, _n_pts) _cs = np.linspace(0.0, 1.0, _n_lines) pts_h = np.stack([np.repeat(_cs, _n_pts), np.tile(_t, _n_lines)], axis=1) pts_v = np.stack([np.tile(_t, _n_lines), np.repeat(_cs, _n_pts)], axis=1) pts_all = jnp.array(np.concatenate([pts_h, pts_v], axis=0)) print(f" Evaluating {len(pts_all)} points for mesh plot ...") vals = jax.device_get(fn(pts_all)) vals_h = vals[: _n_lines * _n_pts].reshape(_n_lines, _n_pts, 2) vals_v = vals[_n_lines * _n_pts :].reshape(_n_lines, _n_pts, 2) for i in range(_n_lines): ax.plot(vals_h[i, :, 0], vals_h[i, :, 1], color=color_h, lw=0.6, zorder=2) ax.plot(vals_v[i, :, 0], vals_v[i, :, 1], color=color_v, lw=0.6, zorder=2) ax.set_xlim(-0.05, 1.05) ax.set_ylim(-0.05, 1.05) ax.set_title(title) ax.set_aspect("equal") fig, axes = plt.subplots(2, 3, figsize=(15, 10)) fig.suptitle("Density transport — Strong BC (post-processing)", fontsize=13) axes[0, 0].contourf(X0, X1, np.ones_like(X0), levels=5, cmap="Blues") axes[0, 0].set_title(r"Source $f=1$") axes[0, 0].set_aspect("equal") cf = axes[0, 1].contourf(X0, X1, G_grid, levels=30, cmap="Reds") axes[0, 1].contour(X0, X1, G_grid, levels=10, colors="k", linewidths=0.4) fig.colorbar(cf, ax=axes[0, 1]) axes[0, 1].set_title(r"Target $g(y)$") axes[0, 1].set_aspect("equal") _plot_mesh_fn(axes[0, 2], T_forward, "b", "r", r"$T_{\rm post}(x)$") def t_roundtrip(x_batch): return T_backward(T_forward(x_batch)) _plot_mesh_fn(axes[1, 0], t_roundtrip, "green", "orange", r"$T^{-1}(T_{\rm post}(x))$") D0, D1 = TX - X0, TY - X1 for col, D, lbl in [(1, D0, r"$T_0 - x_0$"), (2, D1, r"$T_1 - x_1$")]: vm = float(np.abs(D).max()) cf = axes[1, col].pcolormesh( X0, X1, D, cmap="RdBu_r", vmin=-vm, vmax=vm, shading="gouraud" ) fig.colorbar(cf, ax=axes[1, col]) axes[1, col].set_title(lbl) axes[1, col].set_aspect("equal") fig.tight_layout() plt.show()