"""Locating a point in a mesh: is it right, does it vmap, and what does it cost. ``UnstructuredMesh.find_cell_index`` inverts the map of EVERY cell and keeps the one whose preimage lands in the unit square. That is ``n_cells`` Newton solves per point, and it is what made reading a solution cost eleven times the solve that produced it (the measurement is recorded in the docstring of :mod:`scimba_jax.linear_approximation.meshes.point_location`). :mod:`scimba_jax.linear_approximation.meshes.point_location` replaces the search with three ideas that compose: * the **straight-sided** cell -- the polygon through its corners -- ranks candidates with a few plane tests and no iteration at all; * a **background grid** gives a starting guess by index arithmetic, and adjacency says which cells to try next, so what is examined does not grow with the mesh; * the search sits **outside the gradient**, which is what allows a ``while_loop``: a cell index has no derivative, and only the evaluation that follows is differentiated. Two strategies are offered because they fail differently -- ``"rings"`` cannot get stuck but its front grows, ``"greedy"`` is a fixed cost per step but can stall on a concave domain. Both are checked here. **The ground truth is exact, not a second opinion.** The points are generated by mapping reference points of KNOWN cells, so the answer is known before the search runs; comparing against brute force would only say the two agree. The last case is the JET wall, which is where this matters: 483 graded cells of degree 3 around a divertor notch, i.e. 483 Newton solves per point for the method being replaced. """ import time import jax import jax.numpy as jnp import numpy as np from scimba_jax.linear_approximation.meshes.block_structured_mesh import ( BlockStructuredMesh, ) from scimba_jax.linear_approximation.meshes.point_location import ( BlockCellLocator, CellLocator, ) from scimba_jax.linear_approximation.meshes.unstructured_mesh import UnstructuredMesh from scimba_jax.linear_approximation.quad.gauss_quad import UnitSquareTensorized from scimba_jax.mapping.macro_mesh import macro_mesh_ogrid_disk QUAD_ORDER = 4 POINTS_PER_CELL = 3 MANY_PER_CELL = 20 SEED = 0 def disk(cells=4, order=2): """The O-grid disk: curved cells, graded from the centre outwards.""" return UnstructuredMesh.from_macro_mesh( macro_mesh_ogrid_disk(radius=1.0, inner=0.5, n=cells, order=order), UnitSquareTensorized(dim=2, order=QUAD_ORDER), ) def jet_wall(): """The JET wall, meshed once and cached by the FEM examples. Read rather than re-meshed: GMSH does not repeat itself, and a locator validated on one mesh says nothing about another. """ import pathlib cache = ( pathlib.Path(__file__).resolve().parents[1] / "fem" / "solve" / "classical_approach" / "mesh_jet_p3_h0.13.npz" ) if not cache.exists(): return None stored = np.load(cache) class _Macro: nodes = stored["nodes"] cells = stored["cells"] order = int(round(stored["cells"].shape[1] ** 0.5)) - 1 return UnstructuredMesh.from_macro_mesh( _Macro(), UnitSquareTensorized(dim=2, order=QUAD_ORDER) ) def points_with_known_cells(mesh, per_cell=POINTS_PER_CELL, seed=SEED): """Points whose cell is known BEFORE any search, by construction. Reference points are drawn inside ``[0.1, 0.9]^dim`` and pushed through each cell's own map. ⚠ Kept off the faces on purpose: a point exactly on one belongs to both its cells, and a locator that answered either would be right, so such points cannot test anything. """ generator = np.random.default_rng(seed) reference = generator.uniform( 0.1, 0.9, size=(mesh.n_cells_total, per_cell, mesh.dim) ) cells = jnp.repeat(jnp.arange(mesh.n_cells_total), per_cell) points = jax.vmap(mesh._unit_hypercube_to_cell)( jnp.arange(mesh.n_cells_total), jnp.asarray(reference) ).reshape(-1, mesh.dim) return points, np.asarray(cells), reference.reshape(-1, mesh.dim) def timed(function, argument, repeats=5): """Compile once, then the best of ``repeats`` -- with block_until_ready. ⚠ JAX dispatches asynchronously, so a timer stopped without it measures the dispatch and lets the work leak into whatever is timed next. """ started = time.perf_counter() jax.block_until_ready(function(argument)) compile_seconds = time.perf_counter() - started best = min( ( ( lambda t0: ( jax.block_until_ready(function(argument)), time.perf_counter() - t0, )[1] )(time.perf_counter()) ) for _ in range(repeats) ) return compile_seconds, best def check_unstructured(mesh, label): """Both strategies against the known cells, then the cost against brute force.""" points, truth, reference = points_with_known_cells(mesh) print(f"\n=== {label}: {mesh.n_cells_total} mailles, {len(points)} points") brute = jax.jit(lambda p: mesh.find_cell_index(p)[1]) compile_brute, run_brute = timed(brute, points) exact_brute = float(np.mean(np.asarray(brute(points)) == truth)) print( f" force brute {run_brute * 1e3:8.2f} ms " f"(compile {compile_brute:.1f} s) exact {exact_brute:.3f}" ) for strategy in ("greedy", "rings"): locator = CellLocator(mesh, strategy=strategy) located = jax.jit(locator.locate) compile_seconds, run_seconds = timed(located, points) cells, x_hat, ok = located(points) exact = float(np.mean(np.asarray(cells) == truth)) preimage = float(np.max(np.abs(np.asarray(x_hat) - reference))) print( f" {strategy:10s} {run_seconds * 1e3:12.2f} ms " f"(compile {compile_seconds:.1f} s) exact {exact:.3f} " f"trouve {float(jnp.mean(ok)):.3f} ecart x_hat {preimage:.1e} " f"gain {run_brute / run_seconds:5.1f}x" ) def check_block(cells=4): """The multi-patch container: which patch, then which cell inside it.""" macro = macro_mesh_ogrid_disk(radius=1.0, inner=0.5, n=1, order=2, tol=-1.0) block = BlockStructuredMesh( macro, cells, UnitSquareTensorized(dim=2, order=QUAD_ORDER) ) locator = BlockCellLocator(block) print(f"\n=== disque multipatch: {len(locator.meshes)} patchs de {cells}x{cells}") # Known patch AND cell, by construction, patch by patch. generator = np.random.default_rng(SEED) points, want_patch, want_cell = [], [], [] for patch, mesh in enumerate(locator.meshes): reference = generator.uniform(0.1, 0.9, size=(mesh.n_cells_total, 2)) logical = jax.vmap(mesh._unit_hypercube_to_cell)( jnp.arange(mesh.n_cells_total), jnp.asarray(reference)[:, None, :] )[:, 0] points.append(np.asarray(mesh.mapping.local_mapping(logical))) want_patch.append(np.full(mesh.n_cells_total, patch)) want_cell.append(np.arange(mesh.n_cells_total)) points = jnp.asarray(np.concatenate(points)) want_patch = np.concatenate(want_patch) want_cell = np.concatenate(want_cell) located = jax.jit(locator.locate) compile_seconds, run_seconds = timed(located, points) patch, cell, x_hat, ok = located(points) print( f" {len(points)} points {run_seconds * 1e3:.2f} ms " f"(compile {compile_seconds:.1f} s) " f"patch exact {float(np.mean(np.asarray(patch) == want_patch)):.3f} " f"maille exacte {float(np.mean(np.asarray(cell) == want_cell)):.3f} " f"trouve {float(jnp.mean(ok)):.3f}" ) def check_vmap(mesh): """It must vmap, and vmapping must not change the answer. ⚠ The point of the check: the search is a ``while_loop``, and under ``vmap`` every point runs the SAME number of iterations -- the worst one's. So a batched call is not a loop over single calls, and the two agreeing is what says the masking is right. """ points, _, _ = points_with_known_cells(mesh, per_cell=1) locator = CellLocator(mesh, strategy="greedy") batched = np.asarray(jax.jit(locator.locate)(points)[0]) one_by_one = np.array( [int(locator.locate_one(jnp.asarray(p))[0]) for p in np.asarray(points)[:40]] ) print( f"\n=== vmap: {np.array_equal(batched[:40], one_by_one)} " f"(lot contre appels unitaires, {len(one_by_one)} points)" ) def throughput(mesh, label, per_cell=MANY_PER_CELL): """Many points at once: what the batched call costs per point. ⚠ Where the walk is meant to be used. A single point pays the compile and the worst case; a batch amortises the first and shares the second, and the per-point figure is the only one that says whether reading a solution at its measurement points is affordable. """ points, truth, _ = points_with_known_cells(mesh, per_cell=per_cell) locator = CellLocator(mesh, strategy="greedy") located = jax.jit(locator.locate) compile_seconds, run_seconds = timed(located, points) cells, _, ok = located(points) print( f" {label:16s} {len(points):6d} points {run_seconds * 1e3:8.2f} ms " f"{run_seconds / len(points) * 1e6:6.2f} us/point " f"exact {float(np.mean(np.asarray(cells) == truth)):.4f} " f"trouve {float(jnp.mean(ok)):.4f} (compile {compile_seconds:.1f} s)" ) def throughput_block(cells=8, per_cell=MANY_PER_CELL): """The same, on the multi-patch disk.""" macro = macro_mesh_ogrid_disk(radius=1.0, inner=0.5, n=1, order=2, tol=-1.0) block = BlockStructuredMesh( macro, cells, UnitSquareTensorized(dim=2, order=QUAD_ORDER) ) generator = np.random.default_rng(SEED) points, want_patch, want_cell = [], [], [] for patch, mesh in enumerate(block.meshes): reference = generator.uniform(0.1, 0.9, size=(mesh.n_cells_total * per_cell, 2)) which = np.repeat(np.arange(mesh.n_cells_total), per_cell) logical = jax.vmap(mesh._unit_hypercube_to_cell)( jnp.asarray(which), jnp.asarray(reference)[:, None, :] )[:, 0] points.append(np.asarray(mesh.mapping.local_mapping(logical))) want_patch.append(np.full(len(which), patch)) want_cell.append(which) points = jnp.asarray(np.concatenate(points)) want_patch, want_cell = np.concatenate(want_patch), np.concatenate(want_cell) locator = BlockCellLocator(block) located = jax.jit(locator.locate) compile_seconds, run_seconds = timed(located, points) patch, cell, _, ok = located(points) print( f" {'multipatch ' + str(cells) + 'x' + str(cells):16s} {len(points):6d} points " f"{run_seconds * 1e3:8.2f} ms {run_seconds / len(points) * 1e6:6.2f} us/point " f"patch {float(np.mean(np.asarray(patch) == want_patch)):.4f} " f"maille {float(np.mean(np.asarray(cell) == want_cell)):.4f} " f"(compile {compile_seconds:.1f} s)" ) def main(): small = disk(cells=4) check_unstructured(small, "disque") check_vmap(small) check_block() bigger = disk(cells=10) check_unstructured(bigger, "disque raffine") wall = jet_wall() if wall is None: print("\n(paroi JET absente du cache, cas saute)") else: check_unstructured(wall, "paroi JET") print("\n=== debit sur de gros lots") throughput_block() throughput(bigger, "disque 500") if wall is not None: throughput(wall, "paroi JET") if __name__ == "__main__": main()