r"""DG solver for Monge-Ampère via Picard iteration on [0,1]². Solves g(∇u) det(∇²u) = f via the Picard linearisation (Bahari et al.): -Δu^{n+1} + εu^{n+1} = -G(u^n) ∇u^{n+1}·n = x·n on ∂Ω where G(v) = √((Δv)² + f/g(∇v) − det(∇²v)). The DG weak form puts G(u_frozen) in the bilinear form using ``u.freeze("dofsl")``. The Newton solver's jacrev then sees only the linear Laplacian+mass, and each Newton step is one Picard iteration. """ import time import jax import jax.numpy as jnp import matplotlib.pyplot as plt from scimba_jax.linear_approximation.basis.analytic_bases import ( local_lagrange_basis, local_taylor_basis, ) from scimba_jax.linear_approximation.basis.general_bases import AnalyticBasis from scimba_jax.linear_approximation.galerkin.dg.elliptic_dg_scheme import ( EllipticDGscheme, ) from scimba_jax.linear_approximation.galerkin.dg.flux import SIPGFlux from scimba_jax.linear_approximation.galerkin.dg.preconditioner import ( BlockJacobiPreconditioner, ) from scimba_jax.linear_approximation.meshes.mesh import Mesh from scimba_jax.linear_approximation.quad.gauss_quad import UnitSquareTensorized from scimba_jax.linear_approximation.variables.variables_dg import VariablesDG from scimba_jax.mapping.mapping import InvertibleFunction, Mapping from scimba_jax.physical_models.abstract_physical_weak_model import ( AbstractPhysicalWeakModel, ) from scimba_jax.physical_models.abstract_weak_form import AbstractWeakForm # ── Configuration ───────────────────────────────────────────────────────────── N_CELLS = 40 POLY_ORDER = 2 QUAD_ORDER = 4 N_PICARD = 5 EPS_REG = 1e-5 # "taylor" → Taylor basis # "lagrange" → Lagrange basis BASIS_TYPE = "taylor" # "ring" → Gaussian ring # "spiral" → tight spiral centred at (0.7, 0.5) TEST_CASE = "ring" # ── Densities ───────────────────────────────────────────────────────────────── def f_density(x): return 1.0 if TEST_CASE == "ring": def _g_unnorm(y): return 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 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(_g_unnorm)(_pts_zg))) def g_density(y): return _g_unnorm(y) / Z_g # ── Neumann SIPG Flux ───────────────────────────────────────────────────────── class NeumannSIPGFlux(SIPGFlux): """SIPG at interior faces, Neumann ∇u·n = x·n at boundary faces. The ``dirichlet_bc`` function is expected to return the physical coordinate ``x`` itself; the Neumann data g_N = x·n is computed inside ``boundary_call``. """ def boundary_call(self, var_interior, n_interior, bc_val): # bc_val = x (physical coord), n_interior = outward normal g_N = bc_val @ n_interior _u, v, _gradu, _gradv, _fields = var_interior return -g_N * v # ── Picard Monge-Ampère weak form ───────────────────────────────────────────── class MongeAmperePicardWeakForm(AbstractWeakForm): r"""Picard weak form for DG: a(u, v) = ∫ ∇u·∇v + ε u v + G(u*) v where u* = u.freeze("dofsl") and G = √((Δu*)² + f/g(∇u*) − det(∇²u*)). """ eps_reg: float = EPS_REG g_fn: object f_fn: object def __init__(self, dim, f_fn, g_fn, eps_reg=1e-6): super().__init__(dim=dim) self.A = lambda x: jnp.eye(dim) self.f_fn = f_fn self.g_fn = g_fn self.eps_reg = eps_reg def bilinear_form(self, u, v): grad_u = u.gradient("x") grad_v = v.gradient("x") diffusion = grad_v.dot(grad_u) reaction = self.eps_reg * u * v # Frozen G(u*) — no gradient through DOFs u_frozen = u.freeze("dofsl") lap_frozen = u_frozen.laplacian("x") det_H_frozen = u_frozen.det_hessian("x") grad_u_frozen = u_frozen.gradient("x") _g_fn = self.g_fn g_of_grad = _g_fn << grad_u_frozen _f_fn = self.f_fn f_over_g = _f_fn / g_of_grad G_sq = lap_frozen * lap_frozen + f_over_g - det_H_frozen G = G_sq.compose_post_processing_with( lambda val, *a: jnp.sqrt(jnp.clip(val, 1e-12)) ) return diffusion + reaction + G * v def linear_form(self, v): return 0.0 * v # ── Mesh, basis, scheme ─────────────────────────────────────────────────────── physical_dim = 2 out_dim = 1 nb_basis = (POLY_ORDER + 1) ** physical_dim mapping_id = InvertibleFunction(lambda x: x, lambda y: y) mesh = Mesh( dim=physical_dim, n_cells=(N_CELLS, N_CELLS), ref_quad=UnitSquareTensorized(dim=physical_dim, order=QUAD_ORDER), mapping=Mapping(mappings=[mapping_id]), ) _local_basis_fn = local_taylor_basis if BASIS_TYPE == "taylor" else local_lagrange_basis basis = AnalyticBasis( nb_basis=nb_basis, out_dim=out_dim, mesh=mesh, local_basis=lambda coords, i, mesh: _local_basis_fn( coords, i, mesh, order=POLY_ORDER, out_dim=out_dim ), basis_type="scalar", ) variables = VariablesDG(basis=basis, nb_variables=out_dim) pde = MongeAmperePicardWeakForm( dim=physical_dim, f_fn=f_density, g_fn=g_density, eps_reg=EPS_REG ) h = 1.0 / N_CELLS flux = NeumannSIPGFlux(sigma=3, h=h) def neumann_bc(x): return x assembler = EllipticDGscheme( AbstractPhysicalWeakModel.from_weak_form(pde, dirichlet=neumann_bc), variables, flux, ) # ── Initialize DOFs to u = ½|x|² ───────────────────────────────────────────── print("Initializing DOFs to u = ½|x|² ...") assembler.variables.project(lambda x: 0.5 * jnp.sum(x**2, axis=-1, keepdims=True)) # ── Neumann BC function: returns x (physical coordinate) ───────────────────── # ── Picard iteration via Newton solver ──────────────────────────────────────── print(f"Running {N_PICARD} Picard iterations ...") t0 = time.perf_counter() assembler = EllipticDGscheme.solve( assembler, dofsl_init=assembler.variables.dofsl, max_iter=N_PICARD, matrix_free=True, tol=1e-5, preconditioner=BlockJacobiPreconditioner(linear_pde=True), verbose=True, ) jax.block_until_ready(assembler.variables.dofsl) print(f"Done. ({time.perf_counter() - t0:.1f}s)") # ── Evaluation & plots (same pattern as solve_2d_laplacian) ────────────────── N_PLOT = 60 xs = jnp.linspace(0.01, 0.99, N_PLOT) XX, YY = jnp.meshgrid(xs, xs) pts = jnp.stack([XX.ravel(), YY.ravel()], axis=-1) @jax.jit def _eval_u_and_grad(variables_pytree, pts): """Evaluate u_h and ∇u_h on arbitrary points (vectorised, no for loop).""" def _local_eval(pt): _, cell_idx = variables_pytree.mesh.find_cell_index(pt[jnp.newaxis, :]) theta = variables_pytree.dofsl[cell_idx] b = variables_pytree.trial_basis(cell_idx, pt[jnp.newaxis, :]) u = jnp.einsum("iv,qiv->qv", theta, b)[0] db = variables_pytree.trial_basis.derivative(cell_idx, pt[jnp.newaxis, :]) grad_u = jnp.einsum("iv,qivd->qvd", theta, db)[0] return u, grad_u return jax.vmap(_local_eval)(pts) print("Evaluating on grid ...") u_all, grad_all = _eval_u_and_grad(assembler.variables, pts) jax.block_until_ready((u_all, grad_all)) print("Done.") UH = u_all[:, 0].reshape(N_PLOT, N_PLOT) TX = grad_all[:, 0, 0].reshape(N_PLOT, N_PLOT) TY = grad_all[:, 0, 1].reshape(N_PLOT, N_PLOT) G = jax.jit(jax.vmap(g_density))(pts).reshape(N_PLOT, N_PLOT) D = jnp.sqrt((TX - XX) ** 2 + (TY - YY) ** 2) print("Done.") n_lines = 30 line_idx = jnp.linspace(0, N_PLOT - 1, n_lines, dtype=int) # ── Figure ──────────────────────────────────────────────────────────────────── print("Plotting ...") fig, axes = plt.subplots(1, 4, figsize=(20, 5)) fig.suptitle( f"DG Monge-Ampère Picard — {N_CELLS}×{N_CELLS} Q{POLY_ORDER}, " f"{N_PICARD} iters ({TEST_CASE})", fontsize=13, ) im = axes[0].pcolormesh(XX, YY, G, shading="auto", cmap="Reds") fig.colorbar(im, ax=axes[0]) axes[0].set_title(r"Target $g(y)$") axes[0].set_aspect("equal") im = axes[1].pcolormesh(XX, YY, UH, shading="auto", cmap="viridis") fig.colorbar(im, ax=axes[1]) axes[1].set_title(r"Potential $u(x)$") axes[1].set_aspect("equal") for i in line_idx: axes[2].plot(TX[i, :], TY[i, :], "b", lw=0.4) axes[2].plot(TX[:, i], TY[:, i], "r", lw=0.4) axes[2].set_xlim(-0.05, 1.05) axes[2].set_ylim(-0.05, 1.05) axes[2].set_title(r"$T(x) = \nabla u(x)$") axes[2].set_aspect("equal") im = axes[3].pcolormesh(XX, YY, D, shading="auto", cmap="turbo") fig.colorbar(im, ax=axes[3]) axes[3].set_title(r"$|T(x) - x|$") axes[3].set_aspect("equal") fig.tight_layout() plt.show()