r"""Density-to-density transport on the unit disk — weak BC (Neumann). Example 5.6 (Circle) from Bahari et al.: .. math:: \rho(x,y) = 1 + 5\,\mathrm{sech}\!\bigl(5\bigl((x-\tfrac{\sqrt3}{2})^2 +(y-\tfrac12)^2-\tfrac{\pi^2}{4}\bigr)\bigr) + 5\,\mathrm{sech}\!\bigl(5\bigl((x+\tfrac{\sqrt3}{2})^2 +(y-\tfrac12)^2-\tfrac{\pi^2}{4}\bigr)\bigr) Domain: unit disk centred at the origin. Phase 1 — Monge-Ampère with weak Neumann BC. Phase 2 — Inverse map v. Assembly — InvertibleNetBasedICNN. Mesh — Compose square→disk (squircle) with learned OT map via Mapping/Mesh. """ 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 Disk2D, Square2D from scimba_jax.linear_approximation.meshes.mesh import Mesh from scimba_jax.linear_approximation.quad.gauss_quad import UnitSquareTensorized from scimba_jax.mapping.mapping import InvertibleFunction, Mapping 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, ) from scimba_jax.physical_models.function_approximator.function_approximator import ( FunctionApproximator, ) # ── Configuration ────────────────────────────────────────────────────────────── N_COLLOC = 8_000 N_BC_COLLOC = 4_000 N_EPOCHS = 800 N_INV_EPOCHS = 500 N_NEWTON = 5 R_DISK = 1.0 # "three_rings" → Eq. (5.8): three Gaussian rings # "sech" → Eq. (5.9) / Example 5.6: two sech rings TEST_CASE = "three_rings" # ── Densities ───────────────────────────────────────────────────────────────── def f(x): return jnp.array([1.0]) if TEST_CASE == "three_rings": def _g_unnorm(y): return jnp.array( [ 1.0 + 5.0 * jnp.exp( -100.0 * jnp.abs((y[0] - 0.45) ** 2 + (y[1] - 0.4) ** 2 - 0.1) ) + 5.0 * jnp.exp(-100.0 * jnp.abs(y[0] ** 2 + y[1] ** 2 - 0.2)) + 5.0 * jnp.exp( -100.0 * jnp.abs((y[0] + 0.45) ** 2 + (y[1] - 0.4) ** 2 - 0.1) ) ] ) elif TEST_CASE == "sech": def _g_unnorm(y): s3h = jnp.sqrt(3.0) / 2.0 r1sq = (y[0] - s3h) ** 2 + (y[1] - 0.5) ** 2 r2sq = (y[0] + s3h) ** 2 + (y[1] - 0.5) ** 2 pi2_4 = jnp.pi**2 / 4.0 return jnp.array( [ 1.0 + 5.0 / jnp.cosh(5.0 * (r1sq - pi2_4)) + 5.0 / jnp.cosh(5.0 * (r2sq - pi2_4)) ] ) else: raise ValueError(f"Unknown TEST_CASE: {TEST_CASE!r}") _key_zg = jax.random.PRNGKey(42) _k1, _k2 = jax.random.split(_key_zg) _n_mc = 200_000 _r_mc = R_DISK * jnp.sqrt(jax.random.uniform(_k1, (_n_mc,))) _theta_mc = 2.0 * jnp.pi * jax.random.uniform(_k2, (_n_mc,)) _pts_zg = jnp.stack([_r_mc * jnp.cos(_theta_mc), _r_mc * jnp.sin(_theta_mc)], axis=1) Z_g = float(jnp.mean(jax.vmap(lambda y: _g_unnorm(y))(_pts_zg))) def g(y): return _g_unnorm(y) / Z_g # ── Domain ──────────────────────────────────────────────────────────────────── key = jax.random.PRNGKey(0) domain_mu: list = [] dx = Disk2D([0.0, 0.0], R_DISK, is_main_domain=True) dx2 = Square2D([(-3.0, 3.0), (-3.0, 3.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=[12] * 5, 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, 10, N_COLLOC) space_u = _proj_u.space print("Pré-entraînement terminé.") # ══════════════════════════════════════════════════════════════════════════════ # PHASE 1 — Monge-Ampère (weak BC) # ══════════════════════════════════════════════════════════════════════════════ print("\n── Phase 1 : Monge-Ampère (weak BC) ────────────────────────────────") model_ma = MongeAmpere2D(dx, f=f, g=g, model_type="x_mu", bc="weak") sampler_ma = TensorizedSampler( [DomainSampler(dx), UniformParametricSampler(domain_mu)], bc=True ) pinn_ma = Projector( model_ma, space_u, sampler_ma, weights={"interior": [1.0], "boundary": [6.0]}, matrix_regularization=1.2e-5, ) key, sample_dict = sampler_ma.sample(key, N_COLLOC, N_BC_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, N_BC_COLLOC) print( f" best loss: {pinn_ma.best_loss['total']:.4e}" f" | {timeit.default_timer() - t0:.1f}s" ) # ══════════════════════════════════════════════════════════════════════════════ # Level set for boundary post-processing # ══════════════════════════════════════════════════════════════════════════════ EPS_POST = 0.05 level_sets_disk = [ lambda x: R_DISK - jnp.sqrt(x[0] ** 2 + x[1] ** 2), ] # ══════════════════════════════════════════════════════════════════════════════ # 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_disk, eps_post=EPS_POST ) sampler_inv = TensorizedSampler( [DomainSampler(dx), UniformParametricSampler(domain_mu)], bc=False ) pinn_inv_proj = Projector(model_inv, space_v, sampler_inv) key, sample_inv = sampler_inv.sample(key, N_COLLOC) print(f" initial loss: {pinn_inv_proj.evaluate_loss(space_v, sample_inv):.4e}") 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 : InvertibleNetBasedICNN(u, v) # ══════════════════════════════════════════════════════════════════════════════ print("\n── Assembly InvertibleNetBasedICNN ─────────────────────────────────") 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_disk, eps_post=EPS_POST, ) T_forward = jax.jit(jax.vmap(icnn_map)) T_backward = jax.jit(jax.vmap(icnn_map.backward)) # ══════════════════════════════════════════════════════════════════════════════ # Mesh via Mapping (square → disk → adapted disk) # ══════════════════════════════════════════════════════════════════════════════ print("\n── Building meshes ─────────────────────────────────────────────────") def _squircle_fwd(x_ref): u = 2.0 * x_ref[0] - 1.0 v = 2.0 * x_ref[1] - 1.0 return R_DISK * jnp.array( [ u * jnp.sqrt(jnp.clip(1.0 - 0.5 * v**2, 1e-12)), v * jnp.sqrt(jnp.clip(1.0 - 0.5 * u**2, 1e-12)), ] ) def _squircle_inv(y_phys): x0 = jnp.array([0.5 * (y_phys[0] / R_DISK + 1.0), 0.5 * (y_phys[1] / R_DISK + 1.0)]) x0 = jnp.clip(x0, 0.01, 0.99) def step(x, _): res = _squircle_fwd(x) - y_phys J = jax.jacobian(_squircle_fwd)(x) return x - jnp.linalg.solve(J, res), None x_opt, _ = jax.lax.scan(step, x0, None, length=10) return x_opt sq_to_disk = InvertibleFunction(_squircle_fwd, _squircle_inv) ot_fn = InvertibleFunction( lambda x: icnn_map(x), lambda y: icnn_map.backward(y), ) N_MESH = 80 ref_quad = UnitSquareTensorized(dim=2, order=2) mesh_ref = Mesh( dim=2, n_cells=(N_MESH, N_MESH), ref_quad=ref_quad, mapping=Mapping([sq_to_disk]), is_identity_mapping=False, ) mesh_ot = Mesh( dim=2, n_cells=(N_MESH, N_MESH), ref_quad=ref_quad, mapping=Mapping([sq_to_disk, ot_fn]), is_identity_mapping=False, ) def _cell_polygons(mesh, n_edge=12): t = jnp.linspace(0.0, 1.0, n_edge, endpoint=False) edges_unit = jnp.concatenate( [ jnp.stack([t, jnp.zeros_like(t)], axis=1), jnp.stack([jnp.ones_like(t), t], axis=1), jnp.stack([1.0 - t, jnp.ones_like(t)], axis=1), jnp.stack([jnp.zeros_like(t), 1.0 - t], axis=1), ], axis=0, ) def cell_boundary(cell_idx): pts_ref = mesh._unit_hypercube_to_cell(cell_idx, edges_unit) return mesh.mapping.local_mapping(pts_ref) return np.array(jax.vmap(cell_boundary)(mesh.cells_idx)) def _plot_mesh(ax, mesh, title="", color="b"): print(f" Plotting mesh: {title} ...") polys = _cell_polygons(mesh) for boundary in polys: poly = np.vstack([boundary, boundary[0]]) ax.plot(poly[:, 0], poly[:, 1], color=color, lw=0.4) ax.set_aspect("equal") ax.set_title(title) # ── Évaluation sur grille ──────────────────────────────────────────────────── N_GRID = 100 theta_grid = np.linspace(0, 2 * np.pi, N_GRID) r_grid = np.linspace(0, R_DISK, N_GRID) R_G, TH_G = np.meshgrid(r_grid, theta_grid) X0 = R_G * np.cos(TH_G) X1 = R_G * np.sin(TH_G) G_grid = np.vectorize(lambda y0, y1: float(g(jnp.array([y0, y1]))[0]))(X0, X1) # ── Plots ───────────────────────────────────────────────────────────────────── fig, axes = plt.subplots(2, 3, figsize=(15, 10)) fig.suptitle(r"Density transport on disk — weak BC — Example 5.6 (Circle)", 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(axes[0, 2], mesh_ref, r"Reference mesh", "b") _plot_mesh(axes[1, 0], mesh_ot, r"Adapted mesh $T(x)$", "b") xy_flat = jnp.array(np.stack([X0.ravel(), X1.ravel()], axis=1)) mask = np.sqrt(X0**2 + X1**2) <= R_DISK T_vals = np.full((N_GRID * N_GRID, 2), np.nan) valid_idx = np.where(mask.ravel())[0] if len(valid_idx) > 0: T_vals[valid_idx] = jax.device_get(T_forward(xy_flat[valid_idx])) TX = T_vals[:, 0].reshape(N_GRID, N_GRID) TY = T_vals[:, 1].reshape(N_GRID, N_GRID) D0, D1 = TX - X0, TY - X1 Dtot = np.sqrt(np.where(mask, D0**2 + D1**2, np.nan)) for col, D, lbl in [(1, D0, r"$T_0 - x_0$"), (2, D1, r"$T_1 - x_1$")]: D_masked = np.where(mask, D, np.nan) vm = float(np.nanmax(np.abs(D_masked))) cf = axes[1, col].pcolormesh( X0, X1, D_masked, 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()