r"""A BATCH of level sets, classified in one traced program. phiFEM selects the cells it assembles on -- interior and cut -- and lists them. That is the natural thing to do and the one thing a traced program cannot use: the list's length depends on the geometry. Measured on a 32x32 mesh, a circle of radius 0.20 leaves 52 active cells and one of radius 0.35 leaves 120, so ``jit`` and ``vmap`` both refuse with ``NonConcreteBooleanIndexError``. This file shows the fixed-shape answer -- one value per cell rather than a list of the cells that count -- and measures what it buys:: the existing classifier, one level set at a time 0.25 s each masks, jit + vmap over the whole batch ~1 ms for 128 ⚠ **The two are the SAME classification**, not an approximation of it: tags, active cells, ghost faces and boundary faces agree index for index, which the test suite pins on three radii. Three things this file is careful about, each of which was measured ------------------------------------------------------------------ ⚠ **Carry the parameters in a LEAF, not in a closure.** A fresh ``lambda`` per geometry recompiles: 0.25 s each against 0.020 s for a ``ScimbaPytree`` holding ``mu``. On 256 geometries that is 64 s of compilation against 5. ⚠ **Avoid a square root in the level set.** ``sqrt(sum(...)) - 1`` and ``sum(...) - 1`` have the same zero set and give the same classification, but the first has an infinite derivative at its centre -- and a round centre usually lands on a mesh vertex. Measured: the gradient is ``nan``. ⚠ **A hard mask has NO gradient.** Selection by threshold is a step function of the level set, so its derivative is zero almost everywhere and nothing raises. Vectorising does not make a level set learnable; a smooth weight does. Run: python batch_of_level_sets.py """ import time import jax import jax.numpy as jnp import matplotlib.pyplot as plt import numpy as np from scimba_jax.linear_approximation.meshes import levelset_masks as masks from scimba_jax.linear_approximation.meshes.levelset_classifier import ( LevelSetClassifier, ) 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.utils.scimba_pytree import ScimbaPytree N_CELLS = 32 N_GEOMETRIES = 128 EPSILON = 0.03 # width of the smooth transition, in the level set's own units class Ellipse(ScimbaPytree): """``sum(((x - c)/ab)^2) - 1``, carried by a leaf. 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. """ centre, axes = self.mu[:2], self.mu[2:] return jnp.sum(((x - centre) / axes) ** 2) - 1.0 def make_mesh(): """The Cartesian background mesh phiFEM cuts through. Returns: the mesh. """ return Mesh( dim=2, n_cells=(N_CELLS, N_CELLS), ref_quad=UnitSquareTensorized(dim=2, order=4), mapping=Mapping(mappings=[InvertibleFunction(lambda x: x, lambda y: y)]), ) def draw_parameters(key): """A family of ellipses: centres and half-axes. Args: key: a random generator state. Returns: ``(N_GEOMETRIES, 4)``. """ keys = jax.random.split(key, 4) return jnp.stack( [ jax.random.uniform(keys[0], (N_GEOMETRIES,), minval=0.42, maxval=0.58), jax.random.uniform(keys[1], (N_GEOMETRIES,), minval=0.42, maxval=0.58), jax.random.uniform(keys[2], (N_GEOMETRIES,), minval=0.15, maxval=0.35), jax.random.uniform(keys[3], (N_GEOMETRIES,), minval=0.15, maxval=0.35), ], axis=-1, ) def main(): """Classify a family of geometries, and measure what it costs.""" mesh = make_mesh() vertices = masks.cell_vertices(mesh) # geometry: computed ONCE parameters = draw_parameters(jax.random.PRNGKey(0)) def selections(mu): """Everything phiFEM needs, for one geometry, in fixed shape. Args: mu: the ellipse's parameters. Returns: ``(active cells, ghost faces, boundary faces)``. """ tags = masks.cell_tags(Ellipse(mu), vertices) neighbours = masks.face_tags_from_cells(mesh, tags) return ( masks.active_mask(tags), masks.ghost_face_mask(neighbours), masks.boundary_face_mask(neighbours), ) # ── The batch ──────────────────────────────────────────────────────────── batched = jax.jit(jax.vmap(selections)) start = time.perf_counter() active, ghost, boundary = jax.block_until_ready(batched(parameters)) compile_seconds = time.perf_counter() - start start = time.perf_counter() active, ghost, boundary = jax.block_until_ready(batched(parameters)) run_seconds = time.perf_counter() - start print(f"{N_GEOMETRIES} geometries on a {N_CELLS}x{N_CELLS} mesh") print(f" trace + compile : {compile_seconds:.2f} s") print(f" then : {run_seconds * 1000:.2f} ms for the whole batch") counts = active.sum(axis=1) print( f" active cells : {int(counts.min())} to {int(counts.max())} " f"-- the COUNT varies, the SHAPE does not ({active.shape})" ) # ── The witness: one at a time, the old way ────────────────────────────── # ⚠ Only a handful, and with a fresh lambda each time, which is what one # writes spontaneously -- and what recompiles. sample = 8 start = time.perf_counter() for mu in parameters[:sample]: LevelSetClassifier( mesh, lambda x, m=mu: jnp.sum(((x - m[:2]) / m[2:]) ** 2) - 1.0 ) one_by_one = (time.perf_counter() - start) / sample print( f"\n existing classifier, one at a time : {one_by_one * 1000:.0f} ms each" f" -> {one_by_one * N_GEOMETRIES:.1f} s for {N_GEOMETRIES}" ) print( f" speed-up on the batch : {one_by_one * N_GEOMETRIES / run_seconds:.0f}x" ) # ── The gradient, which is the point of the smooth weight ──────────────── def hard_area(mu): return selections(mu)[0].sum() def smooth_area(mu): return masks.smooth_active_weight(Ellipse(mu), vertices, EPSILON).sum() mu0 = jnp.array([0.5, 0.5, 0.30, 0.30]) print("\nd(active area)/d(cx, cy, ax, ay) at a centred circle of radius 0.3") print( f" hard mask : {np.asarray(jax.grad(hard_area)(mu0)).round(2)}" " <- zero, and nothing raised" ) print(f" smooth : {np.asarray(jax.grad(smooth_area)(mu0)).round(1)}") print(" ⚠ the centre derivatives are zero by SYMMETRY here, not by failure:") off = jnp.array([0.46, 0.53, 0.30, 0.24]) print(f" off-centre : {np.asarray(jax.grad(smooth_area)(off)).round(1)}") # ── Figure ─────────────────────────────────────────────────────────────── figure, axes = plt.subplots(2, 4, figsize=(15, 7.5), constrained_layout=True) grid = np.asarray(active).reshape(N_GEOMETRIES, N_CELLS, N_CELLS) for column in range(4): index = column * (N_GEOMETRIES // 4) axes[0, column].imshow(grid[index].T, origin="lower", cmap="Blues") mu = np.asarray(parameters[index]) axes[0, column].set_title( f"active cells ({int(counts[index])})\n" f"c=({mu[0]:.2f},{mu[1]:.2f}) a=({mu[2]:.2f},{mu[3]:.2f})", fontsize=9, ) weight = np.asarray( masks.smooth_active_weight(Ellipse(parameters[index]), vertices, EPSILON) ).reshape(N_CELLS, N_CELLS) image = axes[1, column].imshow(weight.T, origin="lower", cmap="Blues") figure.colorbar(image, ax=axes[1, column]) axes[1, column].set_title("smooth weight (differentiable)", fontsize=9) for axis in axes.ravel(): 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()