r"""Density transport via composition of two OT maps: T = ∇u₂ ∘ ∇u₁. Phase 1: Train u₁ (ICNN) on g(∇u₁) det(∇²u₁) = f. Phase 2: Freeze u₁, train u₂ (ICNN) on the composed residual: g(∇u₂(∇u₁(x))) · det(∇²u₂(∇u₁(x))) · det(∇²u₁(x)) − f(x) = 0 Phase 3: Learn inverse v(T(x)) = x where T = ∇u₂ ∘ ∇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.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 ( InverseFlowResidual, MongeAmpere2D, boundary_post_processing, ) from scimba_jax.physical_models.function_approximator.function_approximator import ( FunctionApproximator, ) # ── Configuration ───────────────────────────────────────────────────────────── # "ring" → Gaussian ring (default) # "spiral" → tight spiral centred at (0.7, 0.5), Eq. (5.7) of Bahari et al. TEST_CASE = "spiral" N_COLLOC = 7_000 N_BC_COLLOC = 4_000 N_EPOCHS_1 = 500 N_EPOCHS_2 = 500 N_INV_EPOCHS = 500 N_NEWTON = 5 EPS_POST = 0.01 # ── 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]) if TEST_CASE == "ring": 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) ) ] ) elif TEST_CASE == "spiral": def _g_unnorm(y): r = jnp.sqrt((y[0] - 0.7) ** 2 + (y[1] - 0.5) ** 2) theta = jnp.arctan2(y[1] - 0.5, y[0] - 0.7) return jnp.array( [1.0 + 9.0 / (1.0 + (10.0 * r * jnp.cos(theta - 20.0 * r)) ** 2)] ) else: raise ValueError(f"Unknown TEST_CASE: {TEST_CASE!r}") _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 # ── Composed MA residual ───────────────────────────────────────────────────── class MongeAmpereComposedResidual(InteriorResidual): r"""Residual for the second map u₂ given frozen u₁. g(∇u₂(∇u₁(x))) · det(∇²u₂(∇u₁(x))) · det(∇²u₁(x)) − f(x) = 0 """ f: object g: object def __init__(self, domain, f, g, space_u1, model_type="x_mu"): super().__init__(domain=domain, size=1, model_type=model_type) self.f = f self.g = g _space_u1 = space_u1 _grad_u1_eval = ( space_u1.create_variables()[0].gradient("x").vmap_on_physical_variables() ) self._grad_u1_fn = lambda x: _grad_u1_eval( jax.lax.stop_gradient(_space_u1), x[None], jnp.zeros((1, 0)) )[0] def _det_hess_u1(x): J = jax.jacobian(self._grad_u1_fn)(x) return jnp.linalg.det(J) self._det_hess_u1 = _det_hess_u1 def construct_residual(self, *vars: PARAM_FUNC_TYPE) -> PARAM_FUNC_TYPE: u2 = vars[0] from scimba_jax.nonlinear_approximation.model_class.funcparam_vectorial import ( ParamFieldFunction, ParamScalarFunction, ) _grad_u1 = self._grad_u1_fn _det_hess_u1 = self._det_hess_u1 grad_u2 = u2.gradient("x") # T_composed(x) = ∇u₂(∇u₁(x)) t_composed = ParamFieldFunction( u2.dims, lambda *args: grad_u2(args[0], _grad_u1(args[1]), *args[2:]), u2.f_type, "x", ) g_of_t = self.g << t_composed # det(J_{T₂∘T₁}(x)) = det(∇²u₂(∇u₁(x))) · det(∇²u₁(x)) def det_composed_fn(*args): x = args[1] y = _grad_u1(x) # det(∇²u₂(y)) J2 = jax.jacobian(lambda yy: grad_u2(args[0], yy, *args[2:]))(y) det2 = jnp.linalg.det(J2) det1 = _det_hess_u1(x) return jnp.array([det1 * det2]) det_composed = ParamScalarFunction(u2.dims, det_composed_fn, u2.f_type) return g_of_t * det_composed - self.f class MongeAmpereComposed(AbstractPhysicalModel): def __init__(self, main_domain, f, g, space_u1, model_type="x_mu", bc="weak"): super().__init__(main_domain=main_domain) self.physical_residuals: PHYSICAL_RESIDUALS_TYPE = { self.main_domain.get_label(): MongeAmpereComposedResidual( domain=main_domain, f=f, g=g, space_u1=space_u1, model_type=model_type ), } if bc == "weak": from scimba_jax.physical_models.elliptic_pde.monge_ampere import ( MongeAmpereNeumannResidual, ) for boundary in self.boundaries: self.physical_residuals[boundary] = MongeAmpereNeumannResidual( domain=self.boundaries[boundary], model_type=model_type ) # ── Inverse residual for composed map ──────────────────────────────────────── class InverseFlowComposed(AbstractPhysicalModel): """Learn v such that v(T_post(∇u₂(∇u₁(x)))) = x.""" def __init__( self, main_domain, space_u1, space_u2, level_sets=None, eps_post=0.02, model_type="x_mu", ): super().__init__(main_domain=main_domain) _space_u1 = space_u1 _space_u2 = space_u2 _grad_u1_eval = ( space_u1.create_variables()[0].gradient("x").vmap_on_physical_variables() ) _grad_u2_eval = ( space_u2.create_variables()[0].gradient("x").vmap_on_physical_variables() ) _level_sets = level_sets _eps = eps_post def composed_map(*args): x = args[0] y = _grad_u1_eval( jax.lax.stop_gradient(_space_u1), x[None], jnp.zeros((1, 0)) )[0] T_x = _grad_u2_eval( jax.lax.stop_gradient(_space_u2), y[None], jnp.zeros((1, 0)) )[0] if _level_sets is not None: T_x = boundary_post_processing(T_x, x, _level_sets, _eps) return T_x self.physical_residuals: PHYSICAL_RESIDUALS_TYPE = { self.main_domain.get_label(): InverseFlowResidual( domain=main_domain, grad_u_fn=composed_map, 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) # ── Helper: pre-train ICNN to ½|x|² ────────────────────────────────────────── def pretrain_icnn(key, nn, domain): space = ApproximationSpace({"x": 2}, [(nn, "scalar", 1)], model_type="x_mu") key, proj = Projector( FunctionApproximator( main_domain=domain, size=1, model_type="x_mu", f_rhs=lambda *args: jnp.array([0.5 * jnp.dot(args[0], args[0])]), ), space, TensorizedSampler( [DomainSampler(domain), UniformParametricSampler(domain_mu)], bc=False ), ).project(key, space, 10, N_COLLOC) return key, proj.space # ══════════════════════════════════════════════════════════════════════════════ # PHASE 1 — First map u₁: standard Monge-Ampère # ══════════════════════════════════════════════════════════════════════════════ print("\n── Phase 1 : u₁ (standard MA) ──────────────────────────────────────") key, subkey = jax.random.split(key) nn_u1 = ICNN( in_size=2, out_size=1, hidden_sizes=[16] * 3, activation="softplus", key=subkey ) print(" Pre-training u₁ ≈ ½|x|² ...") key, space_u1 = pretrain_icnn(key, nn_u1, dx2) model_ma1 = MongeAmpere2D(dx, f=f, g=g, model_type="x_mu", bc="weak") sampler_ma = TensorizedSampler( [DomainSampler(dx), UniformParametricSampler(domain_mu)], bc=True ) pinn_ma1 = Projector( model_ma1, space_u1, sampler_ma, weights={"interior": [1.0], "boundary": [6.0]}, matrix_regularization=2e-6, ) t0 = timeit.default_timer() key, pinn_ma1 = pinn_ma1.project(key, space_u1, N_EPOCHS_1, N_COLLOC, N_BC_COLLOC) print( f" best loss: {pinn_ma1.best_loss['total']:.4e}" f" | {timeit.default_timer() - t0:.1f}s" ) # ══════════════════════════════════════════════════════════════════════════════ # PHASE 2 — Second map u₂: composed residual (u₁ frozen) # ══════════════════════════════════════════════════════════════════════════════ print("\n── Phase 2 : u₂ (composed MA, u₁ frozen) ──────────────────────────") key, subkey = jax.random.split(key) nn_u2 = ICNN( in_size=2, out_size=1, hidden_sizes=[16] * 3, activation="softplus", key=subkey ) print(" Pre-training u₂ ≈ ½|x|² ...") key, space_u2 = pretrain_icnn(key, nn_u2, dx2) model_ma2 = MongeAmpereComposed( dx, f=f, g=g, space_u1=pinn_ma1.space, model_type="x_mu", bc="weak" ) pinn_ma2 = Projector( model_ma2, space_u2, sampler_ma, weights={"interior": [1.0], "boundary": [6.0]}, matrix_regularization=2e-6, ) t0 = timeit.default_timer() key, pinn_ma2 = pinn_ma2.project(key, space_u2, N_EPOCHS_2, N_COLLOC, N_BC_COLLOC) print( f" best loss: {pinn_ma2.best_loss['total']:.4e}" f" | {timeit.default_timer() - t0:.1f}s" ) # ══════════════════════════════════════════════════════════════════════════════ # PHASE 3 — Inverse map v(T₂(T₁(x))) = x # ══════════════════════════════════════════════════════════════════════════════ print("\n── Phase 3 : inverse v(T₂∘T₁(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 = InverseFlowComposed( dx, space_u1=pinn_ma1.space, space_u2=pinn_ma2.space, level_sets=level_sets_square, eps_post=EPS_POST, ) sampler_inv = TensorizedSampler( [DomainSampler(dx), UniformParametricSampler(domain_mu)], bc=False ) pinn_inv = Projector(model_inv, space_v, sampler_inv) t0 = timeit.default_timer() key, pinn_inv = pinn_inv.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" ) # ══════════════════════════════════════════════════════════════════════════════ # Evaluation & Plots # ══════════════════════════════════════════════════════════════════════════════ print("\n── Evaluation ──────────────────────────────────────────────────────") _u1 = pinn_ma1.space.models[0] _u2 = pinn_ma2.space.models[0] _v = pinn_inv.space.models[0] _grad_u1 = jax.grad(lambda x: _u1(x)[0]) _grad_u2 = jax.grad(lambda x: _u2(x)[0]) def T_raw_single(x): # noqa N802 return _grad_u2(_grad_u1(x)) def T_post_single(x): # noqa N802 return boundary_post_processing(T_raw_single(x), x, level_sets_square, EPS_POST) def T1_single(x): # noqa N802 return _grad_u1(x) T_composed = jax.jit(jax.vmap(T_post_single)) T1_forward = jax.jit(jax.vmap(T1_single)) def T_backward_single(y): # noqa N802 x0 = _v(y) _jac = jax.jacobian(T_post_single) def step(x, _): g_res = T_post_single(x) - y J = _jac(x) return x - jnp.linalg.solve(J, g_res), None x_opt, _ = jax.lax.scan(step, x0, None, length=N_NEWTON) return x_opt T_backward = jax.jit(jax.vmap(T_backward_single)) 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) print(" Computing T₁ ...") T1_vals = jax.device_get(T1_forward(xy_flat)) print(" Computing T₂∘T₁ ...") T_vals = jax.device_get(T_composed(xy_flat)) TX = T_vals[:, 0].reshape(N_GRID, N_GRID) TY = T_vals[:, 1].reshape(N_GRID, N_GRID) # ── Mesh plot ───────────────────────────────────────────────────────────────── _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" Mesh plot: {len(pts_all)} pts ...") 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(r"Composed transport $T = \nabla u_2 \circ \nabla u_1$", 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], T1_forward, "b", "r", r"$T_1 = \nabla u_1$") _plot_mesh_fn(axes[1, 0], T_composed, "b", "r", r"$T_2 \circ T_1$") def t_roundtrip(x_batch): return T_backward(T_composed(x_batch)) print(" Computing roundtrip ...") _plot_mesh_fn(axes[1, 1], t_roundtrip, "green", "orange", r"$T^{-1}(T(x))$") D0, D1 = TX - X0, TY - X1 Dtot = np.sqrt(D0**2 + D1**2) vm = float(Dtot.max()) cf = axes[1, 2].pcolormesh( X0, X1, Dtot, cmap="turbo", vmin=0, vmax=vm, shading="gouraud" ) fig.colorbar(cf, ax=axes[1, 2]) axes[1, 2].set_title(r"$|T(x) - x|$") axes[1, 2].set_aspect("equal") fig.tight_layout() plt.show()