r"""Density-to-density transport — comparison of direct vs Picard MA. Trains two ICNN potentials on the same OT problem using: - **Approach A** — Direct Monge-Ampère: g(∇u) det(∇²u) = f - **Approach B** — Picard-linearised MA: -Δu + εu + G(u*) = 0 For each, an inverse map v is learned and an InvertibleNetBasedICNN is assembled. The final figure compares the two transport meshes and shows the flow difference. """ 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.elliptic_pde.monge_ampere import ( InverseFlowMongeAmpere2D, MongeAmpere2D, MongeAmperePicard2D, ) from scimba_jax.physical_models.function_approximator.function_approximator import ( FunctionApproximator, ) # ── Configuration ────────────────────────────────────────────────────────────── # "ring" → Gaussian ring # "spiral" → tight spiral centred at (0.7, 0.5) TEST_CASE = "spiral" N_COLLOC = 8_000 N_BC_COLLOC = 4_000 N_EPOCHS = 1200 N_INV_EPOCHS = 500 N_NEWTON = 5 # ── 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 # ── Domain & samplers ───────────────────────────────────────────────────────── 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) sampler_ma = TensorizedSampler( [DomainSampler(dx), UniformParametricSampler(domain_mu)], bc=True ) sampler_int = TensorizedSampler( [DomainSampler(dx), UniformParametricSampler(domain_mu)], bc=False ) EPS_POST = 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], ] HIDDEN = [12] * 5 # ── Helper: pretrain ICNN to ½|x|² ─────────────────────────────────────────── def pretrain_icnn(key): key, subkey = jax.random.split(key) nn = ICNN( in_size=2, out_size=1, hidden_sizes=HIDDEN, activation="softplus", key=subkey ) space = ApproximationSpace({"x": 2}, [(nn, "scalar", 1)], model_type="x_mu") key, proj = 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, TensorizedSampler( [DomainSampler(dx2), UniformParametricSampler(domain_mu)], bc=False ), ).project(key, space, 10, N_COLLOC) return key, proj.space # ── Helper: train inverse + assemble ───────────────────────────────────────── def train_inverse_and_assemble(key, space_u, label): print(f"\n── {label} : 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=space_u, level_sets=level_sets_square, eps_post=EPS_POST ) pinn_inv_proj = Projector(model_inv, space_v, sampler_int) 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" ) icnn_map = InvertibleNetBasedICNN( u=space_u.models[0], size=2, v=pinn_inv.space.models[0], n_newton=N_NEWTON, level_sets=level_sets_square, eps_post=EPS_POST, ) T_fwd = jax.jit(jax.vmap(icnn_map)) T_bwd = jax.jit(jax.vmap(icnn_map.backward)) return key, T_fwd, T_bwd, pinn_inv # ══════════════════════════════════════════════════════════════════════════════ # APPROACH A — Direct Monge-Ampère # ══════════════════════════════════════════════════════════════════════════════ print("\n══ Approach A : Direct MA ═══════════════════════════════════════════") print(" Pre-training ICNN A ...") key, space_uA = pretrain_icnn(key) model_A = MongeAmpere2D(dx, f=f, g=g, model_type="x_mu", bc="weak") pinn_A = Projector( model_A, space_uA, sampler_ma, weights={"interior": [1.0], "boundary": [6.0]}, matrix_regularization=1.0e-5, ) t0 = timeit.default_timer() key, pinn_A = pinn_A.project(key, space_uA, N_EPOCHS, N_COLLOC, N_BC_COLLOC) print( f" best loss: {pinn_A.best_loss['total']:.4e}" f" | {timeit.default_timer() - t0:.1f}s" ) key, T_fwd_A, T_bwd_A, pinn_inv_A = train_inverse_and_assemble( key, pinn_A.space, "A (Direct)" ) # ══════════════════════════════════════════════════════════════════════════════ # APPROACH B — Picard Monge-Ampère # ══════════════════════════════════════════════════════════════════════════════ print("\n══ Approach B : Picard MA ═══════════════════════════════════════════") print(" Pre-training ICNN B ...") key, space_uB = pretrain_icnn(key) model_B = MongeAmperePicard2D(dx, f=f, g=g, eps_reg=1e-8, model_type="x_mu", bc="weak") pinn_B = Projector( model_B, space_uB, sampler_ma, weights={"interior": [1.0], "boundary": [6.0]}, matrix_regularization=1.0e-5, ) t0 = timeit.default_timer() key, pinn_B = pinn_B.project(key, space_uB, N_EPOCHS, N_COLLOC, N_BC_COLLOC) print( f" best loss: {pinn_B.best_loss['total']:.4e}" f" | {timeit.default_timer() - t0:.1f}s" ) key, T_fwd_B, T_bwd_B, pinn_inv_B = train_inverse_and_assemble( key, pinn_B.space, "B (Picard)" ) # ══════════════════════════════════════════════════════════════════════════════ # Evaluation # ══════════════════════════════════════════════════════════════════════════════ print("\n── Evaluation ──────────────────────────────────────────────────────") 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_A ...") TA_vals = jax.device_get(T_fwd_A(xy_flat)) print(" Computing T_B ...") TB_vals = jax.device_get(T_fwd_B(xy_flat)) # ── Plots ───────────────────────────────────────────────────────────────────── _n_lines = 60 _n_pts = 300 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: {title} ...") 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.5, zorder=2) ax.plot(vals_v[i, :, 0], vals_v[i, :, 1], color=color_v, lw=0.5, zorder=2) ax.set_xlim(-0.05, 1.05) ax.set_ylim(-0.05, 1.05) ax.set_title(title) ax.set_aspect("equal") def _plot_loss(ax, losses_history, title): for label, hist in losses_history.items(): lw = 1.8 if label == "total" else 1.0 color = "k" if label == "total" else None h = hist.ravel() if (hist.ndim == 1 or hist.shape[1] == 1) else None if h is not None: ax.semilogy(h, label=label, lw=lw, color=color) else: for i in range(hist.shape[1]): ax.semilogy(hist[:, i], label=f"{label}[{i}]", lw=lw) ax.set_title(title) ax.set_xlabel("epoch") ax.set_ylabel("loss") ax.legend(fontsize=7) ax.grid(True, which="both", alpha=0.4) # ── Figure : comparison ────────────────────────────────────────────────────── fig, axes = plt.subplots(2, 4, figsize=(22, 10)) fig.suptitle(f"Direct MA vs Picard MA — {TEST_CASE} density", fontsize=14) # Row 0: target g, potential A, mesh A, mesh B cf = axes[0, 0].contourf(X0, X1, G_grid, levels=30, cmap="Reds") axes[0, 0].contour(X0, X1, G_grid, levels=10, colors="k", linewidths=0.4) fig.colorbar(cf, ax=axes[0, 0]) axes[0, 0].set_title(r"Target $g(y)$") axes[0, 0].set_aspect("equal") _uA = pinn_A.space.models[0] UA_grid = np.vectorize(lambda x0, x1: float(_uA(jnp.array([x0, x1]))[0]))(X0, X1) cf = axes[0, 1].contourf(X0, X1, UA_grid, levels=30, cmap="viridis") axes[0, 1].contour(X0, X1, UA_grid, levels=10, colors="k", linewidths=0.4) fig.colorbar(cf, ax=axes[0, 1]) axes[0, 1].set_title(r"Potential $u_A(x)$ (Direct)") axes[0, 1].set_aspect("equal") _plot_mesh_fn(axes[0, 2], T_fwd_A, "b", "r", r"$T_A$ (Direct MA)") _plot_mesh_fn(axes[0, 3], T_fwd_B, "b", "r", r"$T_B$ (Picard MA)") # Row 1: losses A, losses B, |T_A - T_B|, roundtrip error _plot_loss(axes[1, 0], pinn_A.losses.losses_history, "Loss — Direct MA") _plot_loss(axes[1, 1], pinn_B.losses.losses_history, "Loss — Picard MA") # Flow difference |T_A(x) - T_B(x)| diff = np.sqrt( (TA_vals[:, 0] - TB_vals[:, 0]) ** 2 + (TA_vals[:, 1] - TB_vals[:, 1]) ** 2 ).reshape(N_GRID, N_GRID) cf = axes[1, 2].pcolormesh(X0, X1, diff, cmap="turbo", shading="gouraud") fig.colorbar(cf, ax=axes[1, 2]) axes[1, 2].set_title(r"$|T_A(x) - T_B(x)|$") axes[1, 2].set_aspect("equal") # Roundtrip error for approach A def t_roundtrip_A(x_batch): # noqa N802 return T_bwd_A(T_fwd_A(x_batch)) _plot_mesh_fn(axes[1, 3], t_roundtrip_A, "green", "orange", r"$T_A^{-1}(T_A(x))$") fig.tight_layout() plt.show()