"""Nested mesh hierarchy from a curved boundary: k levels, each 4x the previous. Second GMSH entry point (see :mod:`scimba_jax.mapping.mesh_hierarchy`, next to ``macro_mesh``): give it a **boundary** and one ``mesh_size`` -- that of the **coarsest** level -- and it returns ``n_levels`` MacroMeshes that are exact refinements of one another, plus the fine->coarse cell relation between consecutive levels. This is the geometry side of geometric multigrid. GMSH runs once, on the coarsest level; every finer level splits each quad into 4 *in reference space*. For an order-p Lagrange quad the restriction of the geometry map to a sub-square is still of bidegree p, so a child reproduces its parent exactly -- including a curved boundary edge. The figure draws 4 levels of the same domain, each cell coloured by its **level-0 ancestor** (`turbo`): the colour blocks do not move from panel to panel, which is exactly what "the fine mesh is the coarse mesh cut in 4" means. The script also prints the measurements that back the claim (areas, exact restriction, geometric parenthood, and the boundary error that does *not* improve -- the price of nesting). Needs the optional ``mesh`` extra (pygmsh/meshio/gmsh). Saves ``mesh_hierarchy.png``. """ # %% from pathlib import Path import jax import jax.numpy as jnp import matplotlib.pyplot as plt import numpy as np from scimba_jax.linear_approximation.quad.gauss_quad import UnitSquareTensorized from scimba_jax.mapping.mesh_hierarchy import mesh_hierarchy_from_curve ORDER = 3 # geometry degree: 2=quad9, 3=quad16, 4=quad25, ... N_LEVELS = 4 MESH_SIZE = 0.55 # characteristic length of the COARSEST level only def boundary(t): """A closed, curved, non-circular boundary t -> (x, y), t in [0, 1).""" a = 2 * np.pi * t r = 1.0 + 0.15 * np.cos(3 * a) return r * np.cos(a), 0.75 * r * np.sin(a) def boundary_distance(xy): """Signed-ish distance of physical points to the exact boundary curve.""" ts = np.linspace(0.0, 1.0, 4001, endpoint=False) ref = np.stack(boundary(ts), axis=-1) d = np.linalg.norm(xy[:, None, :] - ref[None, :, :], axis=-1) return d.min(axis=1) # %% Build the hierarchy: ONE call, ONE mesh size. hier = mesh_hierarchy_from_curve( boundary, n_levels=N_LEVELS, order=ORDER, mesh_size=MESH_SIZE ) print(hier) for k in range(hier.n_levels): mm = hier[k] print( f" level {k}: {mm.n_cells:5d} cells, {mm.n_nodes:5d} nodes, " f"{mm.n_interfaces:5d} interfaces, {mm.n_boundary_faces:4d} boundary faces" ) # %% Reference-space samples reused for drawing and for the measurements. _S = jnp.linspace(0.0, 1.0, 30) _ZERO, _ONE = jnp.zeros_like(_S), jnp.ones_like(_S) _REF_EDGES = { 0: jnp.stack([_S, _ZERO], -1), 1: jnp.stack([_ONE, _S], -1), 2: jnp.stack([_S, _ONE], -1), 3: jnp.stack([_ZERO, _S], -1), } _LOOP = jnp.concatenate( [_REF_EDGES[0], _REF_EDGES[1], _REF_EDGES[2][::-1], _REF_EDGES[3][::-1]] ) QUAD = UnitSquareTensorized(2, order=8) def cell_area(mapping): """Area of a cell: int_[0,1]^2 det J, by tensorized Gauss quadrature.""" dets = jax.vmap(lambda p: jnp.linalg.det(jax.jacfwd(mapping.forward)(p)))( QUAD.volumic_points ) return float(jnp.sum(QUAD.volumic_weights * dets)) # %% Measurements -- the claims of the module, re-checked on this domain. # Every mapping is a distinct Python closure, so each `jax.vmap(m.forward)` is a # fresh trace: measuring all 1020 cells costs minutes. A random sample of parents # per transition says the same thing in seconds (the unit tests are exhaustive). N_SAMPLE = 12 print("\n--- exact nesting -----------------------------------------------------") print(f" ({N_SAMPLE} random parents per transition; tests/ checks all of them)") rng = np.random.default_rng(0) probe = jnp.asarray(rng.random((25, 2))) for k in range(hier.n_levels - 1): coarse, fine = hier[k], hier[k + 1] sample = rng.choice( coarse.n_cells, size=min(N_SAMPLE, coarse.n_cells), replace=False ) worst_area, worst_map, worst_inv, wrong_inv = 0.0, 0.0, 0.0, np.inf for c in sample: parent = coarse.cell_mapping(int(c)) children = hier.children_of(k)[c] # the 4 children tile the parent: their areas add up to the parent's area_p = cell_area(parent) area_c = sum(cell_area(fine.cell_mapping(int(j))) for j in children) worst_area = max(worst_area, abs(area_c - area_p) / abs(area_p)) for local, child_id in enumerate(children): a, b = divmod(local, 2) child = fine.cell_mapping(int(child_id)) got = jax.vmap(child.forward)(probe) want = jax.vmap(parent.forward)((probe + jnp.array([a, b], float)) / 2.0) worst_map = max(worst_map, float(jnp.max(jnp.abs(got - want)))) # geometric check of the parenthood: the fine centroid must invert # INSIDE the declared parent (inverse() clamps to [0,1]^2, so a # point outside comes back with a large residual). q = child.forward(jnp.array([0.5, 0.5])) res = jnp.max(jnp.abs(parent.forward(parent.inverse(q)) - q)) worst_inv = max(worst_inv, float(res)) # negative control: the same test against a *wrong* parent other = coarse.cell_mapping(int((c + 1) % coarse.n_cells)) bad = jnp.max(jnp.abs(other.forward(other.inverse(q)) - q)) wrong_inv = min(wrong_inv, float(bad)) print( f" {k}->{k + 1}: sum(child areas) vs parent, max rel = {worst_area:.2e} | " f"max|child(u) - parent((u+s)/2)| = {worst_map:.2e}" ) print( f" fine centroid in its parent: residual {worst_inv:.2e} " f"| WRONG parent (negative control): {wrong_inv:.2e}" ) print(f" node-sharing round-off (seam gap) = {hier.seam_gaps[k]:.2e}") print("\n--- what refinement does NOT improve ----------------------------------") print(" the boundary is resolved once, at level 0, and then only subdivided:") for k in range(hier.n_levels): mm = hier[k] pts = np.concatenate( [ np.asarray(jax.vmap(mm.cell_mapping(int(p)).forward)(_REF_EDGES[int(e)])) for p, e in mm.boundary_faces ] ) print( f" level {k}: max distance to the exact curve = " f"{float(np.max(boundary_distance(pts))):.3e}" ) # %% Plot: one panel per level, coloured by the level-0 ancestor. cmap = plt.get_cmap("turbo") n_coarse = hier[0].n_cells fig, axes = plt.subplots(1, N_LEVELS, figsize=(4.6 * N_LEVELS, 5.0)) for k, ax in enumerate(axes): mm = hier[k] ancestor = hier.ancestor(k, np.arange(mm.n_cells), 0) for i in range(mm.n_cells): loop = np.asarray(jax.vmap(mm.cell_mapping(i).forward)(_LOOP)) col = cmap((ancestor[i] + 0.5) / n_coarse) ax.fill(loop[:, 0], loop[:, 1], color=col, alpha=0.55) ax.plot(loop[:, 0], loop[:, 1], color="k", lw=0.35) for patch, edge in mm.boundary_faces: seg = np.asarray( jax.vmap(mm.cell_mapping(int(patch)).forward)(_REF_EDGES[int(edge)]) ) ax.plot(seg[:, 0], seg[:, 1], "k-", lw=1.6) ax.set_title(f"level {k} — {mm.n_cells} cells (h ≈ {MESH_SIZE / 2**k:.3g})") ax.set_aspect("equal") ax.set_xticks([]) ax.set_yticks([]) fig.suptitle( f"mesh_hierarchy_from_curve(order={ORDER}, mesh_size={MESH_SIZE}) — each level " "is the previous one with every cell split in 4\n" "(colour = level-0 ancestor: the blocks never move)", fontsize=13, ) out = Path(__file__).with_name("mesh_hierarchy.png") fig.tight_layout() fig.savefig(out, dpi=140) print(f"\nSaved figure to {out}") plt.show()