"""Domain decomposition: a polynomial basis in the middle, an enriched one outside. -Delta u = f on the disk of radius 2, u = 0 on the circle u(x, y) = (1 - r^2/4) exp(-2 r^2) the geometry benchmark's disk with its envelope CENTRED, so that the whole variation of the solution sits in the middle patch. Five patches, DG-SIPG on every one of them, and only the basis changes: ============== ============== =============================== =========== patch mesh basis learned? ============== ============== =============================== =========== centre 8x8, 64 cells Lagrange Q2, 9 per cell no 4 petals 1 cell each Q2 PLUS ``b(y) N(x)``, 10 each yes, 4 nets ============== ============== =============================== =========== **The petals' basis is the classical one, enriched**:: phi_1..9(x) = Q2 Lagrange the ordinary basis, untouched phi_10(x) = b(y(x)) * N_p(x) b = the cell's bubble, N = the network and the two halves each answer a failure this file measured before them: * **the Lagrange part carries the coupling.** An earlier version gave the petals the network AND NOTHING ELSE, one basis function and one DOF. In SIPG the interface terms are carried by the TEST function of each side, so a petal whose only basis function goes to zero takes its own interface rows with it: the patch DISCONNECTS, and "petal at zero" becomes a free fixed point. It found it -- measured, zeroing all four petals costs 6.89e-02 and the trained run landed on 6.98e-02, one percent away. Flooring that basis (``0.5 + softplus``) closed the escape and made the error worse (1.38e-01), because it constrained the shape without giving it a reason to move. With nine polynomial functions beside it, the question does not arise: the coupling no longer depends on what the network does. * **the bubble keeps the enrichment inside the cell.** ``b`` vanishes on the cell boundary, so the tenth function contributes nothing to any face term -- neither the interface flux nor the Dirichlet one. The network can only change the interior of a cell, which is exactly what an enrichment should do. ⚠ ``b`` is the cell bubble SQUARED, and that is a rank condition -- see :func:`cell_bubble`. **No warm start.** An earlier version pre-trained the network to be ``1`` through ``FunctionApproximator``, because a petal carrying the network alone starts on a discretisation that solves nothing. Here the Lagrange part solves at epoch 0 whatever the network is, so there is nothing to repair before starting and the pre-training was removed rather than kept as decoration. ⚠ **Dirichlet is imposed BY THE FLUX**, weakly, on every patch alike:: -(grad u.n, v) - (u, grad v.n) + sigma/h (u, v) = -(g, grad v.n) + sigma/h (g, v) No nodal value appears anywhere. A continuous FEM patch would impose it on the DOFs instead, by lifting -- the asymmetry ``CLAUDE.md`` records between the two families, seen from the side where it costs nothing. ⚠ **A ragged container**: 9 basis functions a cell in the centre against 10 in the petals, and 64 cells against 4. Its unknowns cannot be stacked on a patch axis, so it presents them flat -- a path that exists because of this file. ⚠ **What the loss can and cannot see.** The petals are 68 percent of the area, 4.7 percent of the squared source and 0.5 percent of the squared solution. Any optimiser is therefore nearly indifferent to them, which is why the previous design could not be rescued by training harder. Per-label ``weights`` on the ``Projector`` are the untried lever. """ import time import jax import jax.numpy as jnp import matplotlib.pyplot as plt import numpy as np from scimba_jax.domains.meshless_domains.domains_2d import Disk2D from scimba_jax.linear_approximation.basis.analytic_bases import ( local_lagrange_basis, local_lagrange_basis_by_logical, ) from scimba_jax.linear_approximation.basis.general_bases import ( AnalyticBasis, PatchwiseParametricBasis, basis_values, ) from scimba_jax.linear_approximation.galerkin.dg.block_structured_dg_scheme import ( BlockStructuredDGscheme, ) 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.meshes.block_structured_mesh import ( BlockStructuredMesh, physical_boundary_sides, ) from scimba_jax.linear_approximation.quad.gauss_quad import UnitSquareTensorized from scimba_jax.linear_approximation.solvers import LinearSolve from scimba_jax.linear_approximation.variables.variables_dg import VariablesDG from scimba_jax.mapping.macro_mesh import macro_mesh_ogrid_disk from scimba_jax.nonlinear_approximation.approximation_spaces.dg_approximation_spaces import ( # noqa: E501 DGEllipticApproximationSpace, ) from scimba_jax.nonlinear_approximation.integration.monte_carlo import ( DomainSampler, TensorizedSampler, ) from scimba_jax.nonlinear_approximation.networks.mlp import MLP from scimba_jax.nonlinear_approximation.numerical_solvers.projectors import Projector from scimba_jax.physical_models.abstract_physical_weak_model import ( AbstractPhysicalWeakModel, ) from scimba_jax.physical_models.classical_weakform.laplacian_weak_form import ( LaplacianWeakForm, ) from scimba_jax.physical_models.elliptic_pde.laplacians import LaplacianDirichletDG from scimba_jax.physical_models.weak_boundary_conditions import Dirichlet jax.config.update("jax_enable_x64", True) DIM = 2 ORDER = 2 # Q2 in the FEM patch GEOMETRY_ORDER = 4 # the petals are ONE cell, so their outer edge is a 90-degree # arc: at degree 2 it misses the circle by 1.1 percent of the radius and the # whole case sits on a geometry floor. Measured in the sibling file. QUAD_ORDER = 4 RADIUS = 2.0 # 4x4 in the centre, one cell per petal: the ragged container this file is about. BLOCK_CELLS = [8, 2, 2, 2, 2] # One enrichment function per cell, ON TOP of the nine Lagrange ones. N_ENRICH = 1 PETAL_NB_BASIS = (ORDER + 1) ** DIM + N_ENRICH HIDDEN = [10, 10, 10] N_COLLOC = 2000 N_EPOCHS = 200 SEED = 0 # The error is read strictly inside the discrete circle: a point beyond it # belongs to no cell and is extrapolated rather than rejected. ERROR_RADIUS = 0.98 * RADIUS SAMPLING_RADIUS = ERROR_RADIUS # Regularisation of the natural gradient's Gram matrix -- the difference between # a run and a `nan` rather than a tuning knob, for the reason the sibling file # documents at length. DAMPING = 5e-6 # Filled bands and black contour lines of the figure. Kept as constants because # they are the only thing that makes a smooth field READABLE: the bands say # where the value is, the lines say how fast it moves, and on this solution -- # a bump inside the FEM patch and a flat tail over the network ones -- the # gradient is the whole story. FILL_LEVELS = 60 LINE_LEVELS = 30 # ⚠ **One training in two stages, not two trainings.** Adam first and ENG # second, the second STARTING FROM the first -- which is what the costs make # obvious once they are separated (see the sibling file): an Adam epoch is # cheap, an ENG epoch assembles the Jacobian of the residual with respect to # every weight, and on a container most of a short ENG run is compilation # anyway. So Adam is asked to do the travelling and ENG the arriving, where a # natural gradient earns its price: near the solution it rescales by the Gram of # the parametrisation, in the directions plain descent leaves flat. SCHEDULE = (("Adam", 10), ("ENG", 120)) OPTIMIZER_KWARGS = { "ENG": {"matrix_regularization": DAMPING}, "Adam": {"learning_rate": 1e-2}, } # ── The manufactured problem ──────────────────────────────────────────────── def u_exact(x): """``(R^2 - r^2)/4``: the solution of ``-Delta u = 1``, zero on the circle. ⚠ Chosen for being SPREAD, which is what the previous choices were not. The benchmark's ``(1 - r^2/4) exp(x + y/2)`` grows towards the rim, so the four one-cell petals carried the largest part of the solution and could not; the centred ``exp(-2 r^2)`` put everything inside the FEM patch instead, and then the learned part had nothing to gain -- measured, perfect petals would have bought 10 percent of a global error the loss cannot resolve to better than 25 percent. A constant source spreads the solution over the whole disk: 1 at the centre, 0.75 at the corner of the middle patch, 0 on the circle. The petals then hold a share of the SOLUTION and of the ERROR that is worth reducing, which is the only setting in which a learned enrichment can be judged at all. """ return jnp.array([(RADIUS**2 - (x[0] ** 2 + x[1] ** 2)) / 4.0]) def f_source(x): """``-Delta u = 1``, written down rather than differentiated. ⚠ It used to be ``-trace(jacfwd(jacfwd(u_exact)))``, which is right and costs a fortune in GRAPH: measured on the traced residual of this container, that autodiff was 2.9 percent of 32 626 jaxpr equations -- a double forward Jacobian of a quadratic, re-traced at every quadrature point, to produce the constant one. The generic form is worth keeping while the solution is being chosen; once it is fixed, the source is data. """ return jnp.ones(1) def g_zero(x): return jnp.zeros(1) # ── The two kinds of patch ────────────────────────────────────────────────── # ⚠ Built ONCE and shared: a container groups its patches by the treedef of # their scheme, callables live in the aux_data, and a lambda rebuilt per patch # gives each one a group of its own. _LOCAL_BASIS = lambda y, i, m: local_lagrange_basis( # noqa: E731 y, i, m, order=ORDER, out_dim=1 ) _LOCAL_BASIS_BY_LOGICAL = lambda y, i, m: local_lagrange_basis_by_logical( # noqa: E731 y, i, m, order=ORDER, out_dim=1 ) _FLUX = SIPGFlux(sigma=50.0 * ORDER * (ORDER + 1), h=None) _QUAD = UnitSquareTensorized(dim=DIM, order=QUAD_ORDER) def fem_basis(mesh): """The classical patch: Lagrange Q2, nothing learned.""" return AnalyticBasis( nb_basis=(ORDER + 1) ** DIM, out_dim=1, mesh=mesh, local_basis=_LOCAL_BASIS, local_basis_by_logical=_LOCAL_BASIS_BY_LOGICAL, basis_type="scalar", ) def cell_bubble(y): """The cell's level set, squared: zero on its boundary, one at its centre. ⚠ **Squared, and that is a rank condition rather than decoration.** ``4 y (1-y)`` per direction is exactly the Q2 bubble -- the ONE function of Q2 that vanishes on the whole boundary of the cell, the subspace being of dimension one. With the warm start below (``N = 1``) the enrichment would then lie in the span of the nine Lagrange functions, and the local block would be singular BY CONSTRUCTION. Squaring puts it at bidegree 4, outside Q2, so the tenth function is independent from the first epoch. Written as a product rather than with an exponent, following the rule this codebase learned the hard way: never raise to a power under ``vmap``. Args: y: Unit-cell coordinates, ``(dim,)``. Returns: The scalar bubble. """ factors = 4.0 * y * (1.0 - y) return jnp.prod(factors * factors) def _reference_point(mesh, cell, x): """The unit-cell coordinates of a physical point -- the physical path. The same two steps ``local_lagrange_basis`` takes: undo the patch mapping, then the affine cell map. """ logical = mesh.mapping.local_inverse_mapping(x[jnp.newaxis, :])[0] return mesh._cell_to_unit_hypercube(cell, logical) def _enriched(lagrange, u, y): """``[Q2 ; bubble(y) * N]``, the petal basis as ``(nb_basis, out_dim)``.""" return jnp.concatenate( [lagrange, (cell_bubble(y) * u).reshape(N_ENRICH, 1)], axis=0 ) def _network_basis(u, x, i, mesh): """Q2 plus the enrichment, at a PHYSICAL point.""" return _enriched( local_lagrange_basis(x, i, mesh, order=ORDER, out_dim=1), u, _reference_point(mesh, i, x), ) def _network_basis_by_logical(u, y, i, mesh): """The same, with the preimage handed over -- the path the assembly takes. ``y`` is already the unit-cell coordinate the bubble wants, so this form is not merely faster: it is the one where the enrichment is written naturally. """ return _enriched( local_lagrange_basis_by_logical(y, i, mesh, order=ORDER, out_dim=1), u, y ) def pinn_basis(mesh, network): """The PINN patch: one network, and its outputs are the basis functions. ⚠ ``local_basis_by_logical`` is the SAME function, which is not laziness: a curved patch asks every basis for the logical form (it sets ``wants_unit_cell_points``), and here there is nothing to write differently -- the polynomial factor that needs the unit cell is exactly what this basis does not have. The network stays a function of the PHYSICAL point on both paths, as it does for every parametric basis. """ return PatchwiseParametricBasis( nb_basis=PETAL_NB_BASIS, out_dim=1, mesh=mesh, patchwise_parametric_function=network, local_basis=_network_basis, local_basis_by_logical=_network_basis_by_logical, basis_type="scalar", ) # ── The warm start ────────────────────────────────────────────────────────── def copies(network, count): """``count`` independent copies of one network, sharing its treedef. ⚠ ``tree_map`` rather than ``MLP(...)`` ``count`` times, and not a shared object either -- see the module docstring. New arrays, same aux_data. ⚠ ``jnp.array``, PAS ``jnp.asarray`` : sur un tableau deja jnp, ``jnp.asarray(a) is a`` vaut ``True``, donc les quatre "copies" portaient litteralement les MEMES feuilles. Personne ne le voyait tant qu'aucune partition ne dedupliquait ; ``auto_partition`` deduplique les feuilles actives par identite, et les fondait donc en une seule -- 261 parametres hors du jit, 1044 dedans, et la Gram levait ``add got incompatible shapes (1044,), (261,)``. """ return [jax.tree_util.tree_map(jnp.array, network) for _ in range(count)] # ── The container ─────────────────────────────────────────────────────────── def _model(sides): """One Laplacian for every patch; ``boundary_scale`` says whose side is whose.""" model = AbstractPhysicalWeakModel(dim=DIM) model.add_weak_form("main", LaplacianWeakForm(dim=DIM, f=f_source)) for side in sides: model.add_boundary_condition(side, Dirichlet(g_zero)) return model def container(petal_basis_of): """The 5-patch disk: FEM in the centre, ``petal_basis_of`` on the rim. Args: petal_basis_of: ``(mesh, petal index) -> basis`` for the four petals. Passing :func:`fem_basis` gives the all-classical reference. Returns: The assembled :class:`BlockStructuredDGscheme`. """ macro = macro_mesh_ogrid_disk( radius=RADIUS, inner=RADIUS / 2, n=1, order=GEOMETRY_ORDER, tol=-1.0 ) block = BlockStructuredMesh(macro, BLOCK_CELLS, _QUAD) sides = [physical_boundary_sides(macro, p) for p in range(macro.n_cells)] every = tuple(sorted({side for patch in sides for side in patch})) shared = _model(every) schemes = [] for patch, mesh in enumerate(block.meshes): basis = fem_basis(mesh) if patch == 0 else petal_basis_of(mesh, patch - 1) schemes.append( EllipticDGscheme( shared, VariablesDG(basis=basis, nb_variables=1), _FLUX, use_scan_quad=False, boundary_scale=[1.0 if s in sides[patch] else 0.0 for s in every], ) ) return BlockStructuredDGscheme(block, schemes) def make_space(scheme): """The differentiable solve. The problem is linear, so it is solved linearly.""" return DGEllipticApproximationSpace( dims={"x": DIM, "dofsl": 1}, list_assemblers=[scheme], model_type="x_dofsl", newton_kwargs={"solver": LinearSolve(tol=1e-8)}, ) def solution_at(space, points): """``u_h`` at a batch of points -- the container locates the patch itself.""" (u_fn,) = space.create_variables() (dofsl,) = space.get_intermediate_values() return jax.vmap(u_fn, in_axes=(None, 0, None))(space, points, dofsl)[:, 0] def error_points(n_points=2000, seed=1): generator = np.random.default_rng(seed) radius = ERROR_RADIUS * np.sqrt(generator.uniform(size=n_points)) angle = generator.uniform(0.0, 2.0 * np.pi, size=n_points) return jnp.asarray(np.stack([radius * np.cos(angle), radius * np.sin(angle)], 1)) def relative_error(space, points, reference): return float( jnp.linalg.norm(solution_at(space, points) - reference) / jnp.linalg.norm(reference) ) # ── Training ──────────────────────────────────────────────────────────────── def train(space, key, schedule=None): """Run the schedule, each stage starting where the previous one stopped. ⚠ Nothing here says which patch is which. ``assemblers_basis_partition`` keeps the parametric functions and freezes the rest, and the centre's basis is analytic -- it carries none, so it is frozen by the same rule that keeps the meshes out. A "learn patches 1 to 4" flag would have been a list where a rule was enough. ⚠ A fresh ``Projector`` per stage, on the PREVIOUS stage's space. An optimiser carries state (Adam's moments, the Gram's damping) that has no meaning across a change of algorithm; what carries over is the space, which is where the networks live. Args: space: The starting approximation space. key: PRNG key. schedule: ``((optimiser, epochs), ...)``, run in order. Returns: ``(projector, history, boundaries, seconds)`` -- ``history`` is the losses of every stage end to end, ``boundaries`` the epoch each stage starts at, for the plot. """ # ⚠ Read at CALL time, not bound as a default: a default freezes ``SCHEDULE`` # at import, so setting the module attribute afterwards -- which is how one # tries a shorter run -- is silently ignored. It cost two runs here. schedule = SCHEDULE if schedule is None else schedule domain = Disk2D(center=(0.0, 0.0), radius=SAMPLING_RADIUS, is_main_domain=True) model = LaplacianDirichletDG(main_domain=domain, f_rhs=f_source, bc="strong") sampler = TensorizedSampler([DomainSampler(domain)], bc=False) histories, boundaries, seconds, projector = [], [], 0.0, None for optimizer, epochs in schedule: projector = Projector( model, space, sampler, optimizer=optimizer, # Defaut `auto_partition` : seuls les champs declares bougent. **OPTIMIZER_KWARGS[optimizer], ) started = time.perf_counter() key, projector = projector.project(key, space, epochs, N_COLLOC, verbose=False) elapsed = time.perf_counter() - started stage = np.asarray(projector.losses.losses_history["total"]).reshape(-1) print( f" {optimizer:4s} {epochs:4d} epoques : loss {stage[0]:.3e} -> " f"{stage[-1]:.3e}, {elapsed:.0f} s" ) boundaries.append(sum(len(h) for h in histories)) histories.append(stage) seconds += elapsed space = projector.space return projector, np.concatenate(histories), boundaries, seconds def main(): points = error_points() reference = jax.vmap(u_exact)(points)[:, 0] stages = " puis ".join(f"{n} {name}" for name, n in SCHEDULE) print( f"Disque de rayon {RADIUS:g} : DG Q{ORDER} Lagrange {BLOCK_CELLS[0]}x" f"{BLOCK_CELLS[0]} au centre, petales {BLOCK_CELLS[1]}x{BLOCK_CELLS[1]} " f"enrichis par un reseau {HIDDEN} ({PETAL_NB_BASIS} fonctions par maille, " f"dont {N_ENRICH} apprise), {stages}" ) # The reference: the same five patches, same method, all polynomial. started = time.perf_counter() classical = make_space(container(lambda mesh, _p: fem_basis(mesh))) classical_error = relative_error(classical, points, reference) print( f" tout Lagrange (reference) : " f"{int(classical.assemblers[0].variables.ndof_linear)}" f" ddl, erreur L2 relative {classical_error:.3e}, " f"{time.perf_counter() - started:.1f} s" ) key = jax.random.PRNGKey(SEED) key, net_key = jax.random.split(key) # ⚠ ONE network, copied four times -- see `copies`. Untrained: the Lagrange # part of the basis already solves, so there is nothing to warm up. petals = copies( MLP(in_size=DIM, out_size=N_ENRICH, hidden_sizes=HIDDEN, key=net_key), 4 ) space = make_space(container(lambda mesh, p: pinn_basis(mesh, petals[p]))) n_dofs = int(space.assemblers[0].variables.ndof_linear) untrained = relative_error(space, points, reference) print(f" Lagrange + reseau : {n_dofs} ddl, warm start {untrained:.3e}") projector, history, boundaries, seconds = train(space, key) error = relative_error(projector.space, points, reference) print(f" apprise : erreur L2 relative {error:.3e}, {seconds:.0f} s au total") print("\n" + "=" * 68) print( f"{'discretisation':22s} {'ddl':>5s} {'L2 depart':>11s} " f"{'L2 finale':>11s} {'temps':>8s}" ) print( f"{'tout Lagrange':22s} " f"{int(classical.assemblers[0].variables.ndof_linear):5d} " f"{classical_error:11.3e} {classical_error:11.3e} {'-':>8s}" ) print( f"{'Lagrange + reseau':22s} {n_dofs:5d} {untrained:11.3e} " f"{error:11.3e} {seconds:7.0f}s" ) result = { "space": projector.space, "history": history, "boundaries": boundaries, "label": stages, } plot([result], classical) def to_physical(mesh, cell, unit_points): """Unit-cell coordinates to PHYSICAL ones, on a block-structured patch. ⚠ Two maps, not one. ``_unit_hypercube_to_cell`` lands in the patch's REFERENCE square and the patch mapping carries that to the disk; applying only the first draws the reference grid over the physical field, which looks like a mesh and is not this one -- measured, it reported an error of 2.3 where the scheme is at 1.9e-02. (On an ``UnstructuredMesh`` the two coincide, its nodes being physical already, which is why the sibling example this was taken from needs no second map.) Args: mesh: The patch mesh. cell: Cell index; may be traced. unit_points: Points of ``[0, 1]^dim``, ``(..., dim)``. Returns: The physical points, same shape. """ points = mesh._unit_hypercube_to_cell(cell, unit_points) mapping = getattr(mesh, "mapping", None) return points if mapping is None else mapping.local_mapping(points) def cell_outlines(mesh, n_samples=25): """The four edges of every cell of one patch, in physical space. Sampled THROUGH the cell map rather than drawn corner to corner: the petals' edges are degree-4 arcs here, and joining their corners with straight lines would draw a different mesh from the one being solved on. Taken from ``solve_laplacian_unstructured_2d``, which needed it for the same reason. """ line = jnp.linspace(0.0, 1.0, n_samples) zeros, ones = jnp.zeros_like(line), jnp.ones_like(line) edges = [ jnp.stack([line, zeros], -1), jnp.stack([ones, line], -1), jnp.stack([line, ones], -1), jnp.stack([zeros, line], -1), ] def one_cell(cell): return jnp.stack([to_physical(mesh, cell, e) for e in edges]) return np.asarray(jax.vmap(one_cell)(jnp.arange(mesh.n_cells_total))) def _cell_triangles(n_cells, n_side): """The triangles of ``n_cells`` copies of an ``n_side x n_side`` grid. ⚠ Triangulated WITHIN each cell rather than left to matplotlib. Given the bare cloud it triangulates the convex hull, which on this disk means edges crossing the interface between two patches -- drawing a continuity that a DG solution does not have and that this file exists to look at. """ corner = np.arange(n_side - 1) i, j = np.meshgrid(corner, corner, indexing="ij") bottom_left = (i * n_side + j).ravel() local = np.concatenate( [ np.stack([bottom_left, bottom_left + n_side, bottom_left + 1], axis=1), np.stack( [bottom_left + n_side, bottom_left + n_side + 1, bottom_left + 1], axis=1, ), ] ) offsets = (np.arange(n_cells) * n_side**2)[:, None, None] return (local[None, :, :] + offsets).reshape(-1, 3) def sample_on_mesh(space, n_side=9): """``u_h`` on a reference grid of every cell of every patch. Patch by patch and cell by cell, so nothing is located and nothing is inverted: the grid point IS the preimage, which is what ``basis_values`` takes as its ``x_hat``. That is the only way to evaluate the FEM patch and a network patch by the same line -- one has a polynomial basis on a curved cell, the other has no polynomial at all. Returns: ``(points, triangles, values)`` for ``tricontourf``. """ container = space.assemblers[0] # ⚠ `get_intermediate_values` stacks over ASSEMBLERS, so the container's own # flat vector is entry 0 -- not the tuple element, which is the stack. (stacked,) = space.get_intermediate_values() blocks = container.split(stacked[0]) grid = jnp.stack( jnp.meshgrid( jnp.linspace(0.0, 1.0, n_side), jnp.linspace(0.0, 1.0, n_side), indexing="ij", ), axis=-1, ).reshape(-1, 2) all_points, all_values, offset, triangles = [], [], 0, [] for patch, scheme in enumerate(container.schemes): mesh = scheme.variables.mesh basis = scheme.variables.trial_basis theta_of_cell = blocks[patch] def one_cell(cell, mesh=mesh, basis=basis, theta_of_cell=theta_of_cell): points = to_physical(mesh, cell, grid) theta = theta_of_cell[cell] def value(point, x_hat): shape = basis_values(basis, cell, point, x_hat) return jnp.einsum("iv,qiv->qv", theta, shape)[0, 0] return points, jax.vmap(value)(points, grid) points, values = jax.vmap(one_cell)(jnp.arange(mesh.n_cells_total)) all_points.append(np.asarray(points).reshape(-1, 2)) all_values.append(np.asarray(values).reshape(-1)) triangles.append(_cell_triangles(int(mesh.n_cells_total), n_side) + offset) offset += int(mesh.n_cells_total) * n_side**2 return ( np.concatenate(all_points), np.concatenate(triangles), np.concatenate(all_values), ) def plot(results, classical): """Solutions on top, errors below, both drawn ON the mesh. ``turbo`` with contour lines, and the cells of the five patches over them: the point of the figure is WHERE the error sits, and a scatter of random points hides both the interface and the one-cell petals. """ panels = [(classical, "tout Lagrange")] + [ (result["space"], f"Lagrange + reseau\n({result['label']})") for result in results ] columns = 1 + len(panels) figure, axes = plt.subplots(2, columns, figsize=(4.6 * columns, 9.0)) outlines = [ cell_outlines(scheme.variables.mesh) for scheme in classical.assemblers[0].schemes ] def draw_mesh(axis): for patch in outlines: for cell in patch: for edge in cell: axis.plot(edge[:, 0], edge[:, 1], "-", color="0.25", lw=0.5) axis.set_aspect("equal") axis.set_xticks([]) axis.set_yticks([]) points, triangles, _ = sample_on_mesh(classical) truth = np.asarray(jax.vmap(u_exact)(jnp.asarray(points)))[:, 0] levels = np.linspace(truth.min(), truth.max(), FILL_LEVELS) lines = np.linspace(truth.min(), truth.max(), LINE_LEVELS) filled = axes[0, 0].tricontourf( points[:, 0], points[:, 1], triangles, truth, levels=levels, cmap="turbo" ) axes[0, 0].tricontour( points[:, 0], points[:, 1], triangles, truth, levels=lines, colors="k", linewidths=0.35, ) draw_mesh(axes[0, 0]) axes[0, 0].set_title("u exacte") figure.colorbar(filled, ax=axes[0, 0], fraction=0.046) for column, (space, title) in enumerate(panels, start=1): pts, tris, values = sample_on_mesh(space) exact = np.asarray(jax.vmap(u_exact)(jnp.asarray(pts)))[:, 0] filled = axes[0, column].tricontourf( pts[:, 0], pts[:, 1], tris, values, levels=levels, cmap="turbo", extend="both", ) axes[0, column].tricontour( pts[:, 0], pts[:, 1], tris, values, levels=lines, colors="k", linewidths=0.35, ) draw_mesh(axes[0, column]) axes[0, column].set_title(title) figure.colorbar(filled, ax=axes[0, column], fraction=0.046) error = np.abs(values - exact) relative = np.linalg.norm(values - exact) / np.linalg.norm(exact) filled = axes[1, column].tricontourf( pts[:, 0], pts[:, 1], tris, error, levels=FILL_LEVELS, cmap="turbo" ) axes[1, column].tricontour( pts[:, 0], pts[:, 1], tris, error, levels=LINE_LEVELS // 2, colors="k", linewidths=0.3, alpha=0.6, ) draw_mesh(axes[1, column]) axes[1, column].set_title(f"|u_h - u|, L2 rel = {relative:.2e}") figure.colorbar(filled, ax=axes[1, column], fraction=0.046) loss_axis = axes[1, 0] loss_axis.axis("on") for result in results: loss_axis.semilogy(result["history"], linewidth=1, color="C0") # Where one optimiser hands over to the next. for boundary, (name, _) in zip(result["boundaries"], SCHEDULE): if boundary > 0: loss_axis.axvline(boundary, color="0.5", linestyle="--", linewidth=1) loss_axis.annotate( name, (boundary, result["history"][boundary]), textcoords="offset points", xytext=(4, 6), fontsize=8, ) loss_axis.set_xlabel("epoque") loss_axis.set_ylabel("residu fort") loss_axis.set_title("apprentissage") loss_axis.grid(True, alpha=0.3) figure.suptitle( "Decomposition de domaine (DG-SIPG partout) : Lagrange Q2 4x4 au centre, " "base-reseau (1 ddl) sur chaque petale" ) plt.tight_layout(rect=(0.0, 0.0, 1.0, 0.97)) plt.show() if __name__ == "__main__": main()