r"""phiFEM solved on a BATCH of geometries, parametric and not. The same Poisson-Dirichlet problem, ``-Delta u = 1`` with ``u = 0`` on the boundary, on many domains at once. Two families, because they need two different mechanisms and the difference is worth seeing: * **parametric** -- a circle deforming into an ellipse. The shape parameters are a ``jnp`` LEAF, so the whole family is one batched pytree and the solve is a single ``vmap``; * **non-parametric** -- a circle, a square with rounded corners, a peanut and a star. These are genuinely different FUNCTIONS, not one shape at different coefficients, so they travel through an :class:`~....utils.functional_fields.AbstractFunctionalField`, exactly as ``uq_batched_transport_1d.py`` carries its transport fields. ⚠ **What makes this possible at all**: phiFEM normally assembles on a LIST of active cells, whose length depends on the geometry -- 52 cells for a circle of radius 0.20, 120 for one of radius 0.35 -- so ``vmap`` refuses. Here nothing is selected: every cell is assembled and weighted, zero outside. The shape is fixed, the content is not. ⚠ **The two mechanisms do not cost the same.** A parametric family is free: one compiled program, the parameters are data. The functional-field registry makes every evaluation of ``phi`` a ``lax.switch``, which under ``vmap`` becomes a ``select_n`` that runs EVERY branch -- and ``phi`` is evaluated at every quadrature point of every cell. The registry is the price of heterogeneity, so prefer parameters whenever the family has any. ⚠ **Write the level set without a square root.** ``sum((x-c)^2) - r^2`` and ``|x-c| - r`` have the same zero set and give the same classification, but the second has an infinite derivative at the centre -- and a round centre usually lands on a mesh vertex. Measured: the gradient comes out ``nan``. ⚠ **Volume terms only** -- see :func:`~...phi_fem.masked_poisson.solve_phifem`. The ghost penalty that stabilises small cut cells is not ported yet, so the error here is that of an unstabilised phiFEM: about 2% against the exact solution on a 32x32 mesh, which is enough to show the batching works and is not a claim about the method's accuracy. Run: python batch_of_geometries.py """ import time import jax import jax.numpy as jnp import matplotlib.pyplot as plt import numpy as np from scimba_jax.linear_approximation.basis.analytic_bases import local_lagrange_basis from scimba_jax.linear_approximation.basis.general_bases import AnalyticBasis from scimba_jax.linear_approximation.meshes.mesh import Mesh from scimba_jax.linear_approximation.phi_fem import masked_poisson as phifem from scimba_jax.linear_approximation.quad.gauss_quad import UnitSquareTensorized from scimba_jax.linear_approximation.variables.variables_fe import VariablesFE from scimba_jax.mapping.mapping import InvertibleFunction, Mapping from scimba_jax.utils.functional_fields import make_functional_field_class from scimba_jax.utils.scimba_pytree import ScimbaPytree N_CELLS, ORDER = 32, 1 N_PARAMETRIC = 32 #: ⚠ ONE class, at module level, shared by every level set. Built per geometry #: instead, each shape would land at index 0 of its own registry and the batch #: would silently collapse onto the first one -- the trap that #: `laplacian_2d_disk_geofno.py` documents at length. LevelSetField = make_functional_field_class("phifem_level_set") def source(x): """``f = 1``, so the exact solution on a disk is ``(R^2 - r^2)/4``. Args: x: a point. Returns: the source there. """ del x return jnp.ones(()) class Ellipse(ScimbaPytree): """``sum(((x - c)/ab)^2) - 1``, the parametric family. ⚠ The parameters are a ``jnp`` leaf, which is what makes a batch of these a single batched pytree -- and what avoids recompiling per geometry. Measured on a 32x32 mesh: a fresh lambda per shape costs 0.25 s each against 0.020 s for a stable object carrying its parameters. Args: mu: ``(4,)`` -- centre x, centre y, half-axis x, half-axis y. """ def __init__(self, mu): self.mu = jnp.asarray(mu, dtype=float) def __call__(self, x): """The level set at one point. Args: x: a point, ``(2,)``. Returns: a scalar, negative inside. """ return jnp.sum(((x - self.mu[:2]) / self.mu[2:]) ** 2) - 1.0 # ── The non-parametric family: genuinely different shapes ──────────────────── def circle(x): """A disk of radius 0.30. Args: x: a point. Returns: the level set there. """ return jnp.sum((x - 0.5) ** 2) - 0.30**2 def rounded_square(x): """A square with rounded corners -- a super-ellipse of exponent 4. Args: x: a point. Returns: the level set there. """ return jnp.sum(((x - 0.5) / 0.30) ** 4) - 1.0 def peanut(x): """Two overlapping disks, joined into one domain. Args: x: a point. Returns: the level set there. """ left = jnp.sum((x - jnp.array([0.40, 0.5])) ** 2) - 0.20**2 right = jnp.sum((x - jnp.array([0.60, 0.5])) ** 2) - 0.20**2 return jnp.minimum(left, right) # union: the min of two level sets def clover(x): """Four lobes along the axes -- a disk pinched on its diagonals. ⚠ Two failed attempts before this one, both worth recording. Writing the shape as ``r^2 - (R + a cos(5 theta))^2`` is singular at the centre, where the angle is undefined: measured ``u(centre) = -2.3e-04``, negative, which the maximum principle forbids. Replacing the angle by the polynomial ``Re((x+iy)^5)`` removes the singularity but not the real problem -- a degree-5 term outgrows ``r^2``, so the zero set is UNBOUNDED whatever the coefficient. Measured: the domain covered 52% of the square and ran into its edges at a=22, still 26% at a=6. A degree-4 term with a POSITIVE sign dominates at infinity instead, so the domain stays closed: 19% of the square, touching nothing. Args: x: a point. Returns: the level set there. """ offset = x - 0.5 pinch = (offset[0] * offset[1]) ** 2 return jnp.sum(offset**2) - 0.27**2 + 25.6 * pinch SHAPES = { "circle": circle, "rounded square": rounded_square, "peanut": peanut, "clover": clover, } def make_space(): """The background Cartesian mesh and its Q1 space. Returns: the variables. """ mesh = Mesh( dim=2, n_cells=(N_CELLS, N_CELLS), ref_quad=UnitSquareTensorized(dim=2, order=2 * ORDER + 2), mapping=Mapping(mappings=[InvertibleFunction(lambda x: x, lambda y: y)]), ) basis = AnalyticBasis( nb_basis=(ORDER + 1) ** 2, out_dim=1, mesh=mesh, basis_type="scalar", local_basis=lambda y, i, m: local_lagrange_basis( y, i, m, order=ORDER, out_dim=1 ), ) return VariablesFE(basis=basis, nb_variables=1) def main(): """Solve both families, and check the parametric one against the exact solution.""" variables = make_space() solver = phifem.batched_solver(variables, source) # ── 1. Parametric: circle -> ellipse ───────────────────────────────────── t = jnp.linspace(0.0, 1.0, N_PARAMETRIC) mus = jnp.stack( [ jnp.full(N_PARAMETRIC, 0.5), jnp.full(N_PARAMETRIC, 0.5), 0.28 - 0.10 * t, 0.28 + 0.10 * t, ], axis=-1, ) start = time.perf_counter() dofs = jax.block_until_ready(solver(Ellipse(mus))) compile_seconds = time.perf_counter() - start start = time.perf_counter() dofs = jax.block_until_ready(solver(Ellipse(mus))) run_seconds = time.perf_counter() - start print(f"1. PARAMETRIC -- {N_PARAMETRIC} geometries, circle to ellipse") print(f" trace + compile : {compile_seconds:.2f} s") print(f" then : {run_seconds * 1000:.0f} ms for the batch") # ⚠ The first member is a circle of radius 0.28, whose solution is known. centre = jax.vmap( lambda mu, d: phifem.evaluate(variables, Ellipse(mu), d, jnp.array([0.5, 0.5])) )(mus, dofs)[:, 0] exact = 0.28**2 / 4.0 print( f" u at the centre : {float(centre[0]):.5f} (circle, exact {exact:.5f}, " f"{abs(float(centre[0]) - exact) / exact:.1%}) -> {float(centre[-1]):.5f} (ellipse)" ) # ── 2. Non-parametric: four different shapes ───────────────────────────── fields = [LevelSetField(shape) for shape in SHAPES.values()] batched_field = jax.tree_util.tree_map(lambda *leaves: jnp.stack(leaves), *fields) start = time.perf_counter() shape_dofs = jax.block_until_ready(solver(batched_field)) print( f"\n2. NON-PARAMETRIC -- {len(SHAPES)} different shapes, one functional field" ) print(f" solved in {time.perf_counter() - start:.2f} s {shape_dofs.shape}") print( f" registry: {len(LevelSetField.function_registry)} entries -- every " "evaluation of phi runs them all under vmap" ) for name, field, d in zip(SHAPES, fields, shape_dofs): value = phifem.evaluate(variables, field, d, jnp.array([0.5, 0.5])) print(f" {name:16s} u(centre) = {float(value[0]):.5f}") # ── Figure ─────────────────────────────────────────────────────────────── side = jnp.linspace(0.0, 1.0, 160) gx, gy = jnp.meshgrid(side, side, indexing="ij") points = jnp.stack([gx.ravel(), gy.ravel()], axis=-1) figure, axes = plt.subplots(2, 4, figsize=(15, 7.5), constrained_layout=True) for column, index in enumerate((0, 10, 20, 31)): field = Ellipse(mus[index]) values = jax.vmap( lambda p, f=field, d=dofs[index]: phifem.evaluate(variables, f, d, p) )(points)[:, 0] inside = jax.vmap(field)(points) <= 0.0 image = np.where(np.asarray(inside), np.asarray(values), np.nan) picture = axes[0, column].pcolormesh( np.asarray(gx), np.asarray(gy), image.reshape(160, 160), shading="auto" ) figure.colorbar(picture, ax=axes[0, column]) mu = np.asarray(mus[index]) axes[0, column].set_title(f"a={mu[2]:.2f} b={mu[3]:.2f}", fontsize=9) for column, (name, field) in enumerate(zip(SHAPES, fields)): values = jax.vmap( lambda p, f=field, d=shape_dofs[column]: phifem.evaluate(variables, f, d, p) )(points)[:, 0] inside = jax.vmap(SHAPES[name])(points) <= 0.0 image = np.where(np.asarray(inside), np.asarray(values), np.nan) picture = axes[1, column].pcolormesh( np.asarray(gx), np.asarray(gy), image.reshape(160, 160), shading="auto" ) figure.colorbar(picture, ax=axes[1, column]) axes[1, column].set_title(name, fontsize=9) for axis in axes.ravel(): axis.set_aspect("equal") axis.set_xticks([]) axis.set_yticks([]) output = __file__.replace(".py", ".png") figure.savefig(output, dpi=110) print(f"\nfigure: {output}") plt.show() if __name__ == "__main__": main()