"""phi-FEM on the Galerkin infrastructure: Krylov, gradient in the level set, batch. The same problem as ``solve_dirichlet_circle.py`` -- ``-Delta u = f`` on the disk ``{phi < 0}``, ``u = 0`` on the circle, imposed by the ansatz ``u_h = phi w_h`` -- solved by :class:`~scimba_jax.linear_approximation.phi_fem. phifem_scheme.PhiFEMscheme` instead of the dense reference: 1. the solve is matrix-free (BiCGSTAB on ``jax.linearize`` of the residual; the boundary term makes the operator nonsymmetric, so not CG), with the errors of the reference to rounding -- measured here on 10x10, 40x40 and 80x80 Q1: same L2 and H1 errors, ``w`` equal to 7e-10, the 80x80 solve 0.85 s warm against 11.5 s for the dense assembly plus ``np.linalg.solve``; 2. the radius of the circle is a PARAMETER: ``jax.grad`` of a loss through the solve reaches it, by the implicit differentiation every Galerkin solve carries. The hard masks (which cells are active) have no gradient; the ansatz ``phi w`` and the penalties do, and the finite difference agrees; 3. a batch of radii is ONE compiled program: the scheme is built once and ``with_level_set`` derives the members under ``jax.vmap``. ⚠ The gradient part runs on a 10x10 background mesh, not 12x12: with ``R = sqrt(2)/4`` the circle passes EXACTLY through the vertices ``(3, 3)`` of the 12x12 grid (``(i - 6)^2 + (j - 6)^2 = 144 R^2 = 18``), so the hard masks flip for any perturbation of the radius and the loss is discontinuous there. The implicit gradient is the derivative of the branch the masks select; a finite difference straddles the jump and measures nothing (1.4e+05 against -9.8e+02 measured). A level set through a vertex is a knife edge for every selection-based cut method; away from it the two agree to 1e-4. Exact solution: ``u = phi exp(x) sin(2 pi y)``, ``phi = (x-1/2)^2 + (y-1/2)^2 - R^2``, ``R = sqrt(2)/4``. """ import time import jax import jax.numpy as jnp 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.phifem_scheme import PhiFEMscheme from scimba_jax.linear_approximation.quad.gauss_quad import UnitSquareTensorized from scimba_jax.linear_approximation.solvers import LinearKrylov from scimba_jax.linear_approximation.variables.variables_fe import VariablesFE from scimba_jax.mapping.mapping import InvertibleFunction, Mapping from scimba_jax.physical_models.classical_weakform.laplacian_weak_form import ( LaplacianWeakForm, ) from scimba_jax.utils.scimba_pytree import ScimbaPytree RADIUS = 2**0.5 / 4 N_CELLS = 40 class Circle(ScimbaPytree): """A level set whose radius is a leaf -- hence a parameter of the scheme.""" def __init__(self, radius): self.radius = jnp.asarray(radius) def __call__(self, x): return (x[0] - 0.5) ** 2 + (x[1] - 0.5) ** 2 - self.radius**2 def u_exact(x): return Circle(RADIUS)(x) * jnp.exp(x[0]) * jnp.sin(2 * jnp.pi * x[1]) def source(x): return -jnp.trace(jax.hessian(u_exact)(x)) def make_space(n_cells, order=1): mesh = Mesh( dim=2, n_cells=[n_cells, n_cells], ref_quad=UnitSquareTensorized(dim=2, order=3), mapping=Mapping(mappings=[InvertibleFunction(lambda x: x, lambda y: y)]), ) basis = AnalyticBasis( nb_basis=(order + 1) ** 2, out_dim=1, mesh=mesh, local_basis=lambda c, i, m: local_lagrange_basis( c, i, m, order=order, out_dim=1 ), basis_type="scalar", ) return VariablesFE(basis=basis, nb_variables=1) solver = LinearKrylov(tol=1e-12, cg_solver="bicgstab", max_iter_linear=5000) # 1. The solve -------------------------------------------------------------- scheme = PhiFEMscheme( LaplacianWeakForm(dim=2, f=source), make_space(N_CELLS), Circle(RADIUS) ) started = time.perf_counter() dofsl, report = PhiFEMscheme.solve_pure(scheme, solver=solver) jax.block_until_ready(dofsl) first = time.perf_counter() - started started = time.perf_counter() dofsl, report = PhiFEMscheme.solve_pure(scheme, solver=solver) jax.block_until_ready(dofsl) warm = time.perf_counter() - started scheme.variables.dofsl = dofsl print( f"phi-FEM Q1 {N_CELLS}x{N_CELLS}: {dofsl.size} nodes, BiCGSTAB {int(report.n_linear)} " f"iterations, residual {float(report.residual):.1e}; first call {first:.2f} s " f"(compilation), warm {warm * 1e3:.0f} ms" ) print( f" L2 error {scheme.l2_error(u_exact, relative=True):.3e}, " f"H1 error {scheme.h1_seminorm_error(u_exact, relative=True):.3e} (relative)" ) # 2. The gradient in the radius --------------------------------------------- space = make_space(10) base = PhiFEMscheme(LaplacianWeakForm(dim=2, f=source), space, Circle(RADIUS)) def loss(radius): """``sum w^2`` of the solve on the circle of that radius.""" dofsl, _ = PhiFEMscheme.solve_pure( base.with_level_set(Circle(radius)), solver=solver ) return jnp.sum(dofsl**2) value, gradient = jax.value_and_grad(loss)(RADIUS) eps = 1e-5 finite = (loss(RADIUS + eps) - loss(RADIUS - eps)) / (2 * eps) print( f"d loss / d radius: implicit {float(gradient):.6e}, " f"finite difference {float(finite):.6e}" ) # 3. A batch of radii, one program ------------------------------------------- radii = jnp.array([0.28, 0.31, 0.34, RADIUS]) solve_batch = jax.jit( jax.vmap( lambda r: PhiFEMscheme.solve_pure( base.with_level_set(Circle(r)), solver=solver )[0] ) ) started = time.perf_counter() batch = jax.block_until_ready(solve_batch(radii)) print( f"{len(radii)} radii in one program: {time.perf_counter() - started:.2f} s " f"(compilation included), shapes {batch.shape}" )