"""Shared by the JET and MAST scripts: the same equilibrium by FEM and by PINN. The PINN side is the model and the two phases of ``examples/examples_jax/pinns/ stationary_pdes/elliptic_pdes/grad_shafranov_with_xpoints.py`` -- same physics constants, same optimisers --, with the boundary weight, the network, the collocation counts and the epochs retuned (constants below): at the example's weight of 100 the boundary was under-resolved. The FEM side solves the same model (:class:`~scimba_jax.physical_models.elliptic_pde.grad_shafranov_fem. GradShafranovFullPhysicsWeakModel`) on a Q2 mesh of the same wall, with the same two phases: the linear source, then the full-physics one with the axis and X-points recomputed before every assembly and a Picard iteration. """ import time from pathlib import Path import jax import jax.numpy as jnp import matplotlib.pyplot as plt import numpy as np from matplotlib.path import Path as PolygonPath from matplotlib.tri import Triangulation from scimba_jax.domains.meshless_domains.domains_2d import Square2D from scimba_jax.domains.tokamak import ( TokamakSampler2D, eqdsk_wall_loop, read_eqdsk, read_eqdsk_wall, ) from scimba_jax.linear_approximation.basis.analytic_bases import ( local_lagrange_basis, local_lagrange_basis_by_logical, ) from scimba_jax.linear_approximation.basis.dof_map import UnstructuredLagrangeDofMap from scimba_jax.linear_approximation.basis.general_bases import AnalyticBasis from scimba_jax.linear_approximation.galerkin.fem.elliptic_fe_scheme import ( EllipticFEscheme, ) from scimba_jax.linear_approximation.meshes.unstructured_mesh import UnstructuredMesh from scimba_jax.linear_approximation.quad.gauss_quad import UnitSquareTensorized from scimba_jax.linear_approximation.variables.variables_fe import VariablesFE from scimba_jax.mapping.macro_mesh import macro_mesh_from_points from scimba_jax.nonlinear_approximation.approximation_spaces.approximation_spaces import ( # noqa: E501 ApproximationSpace, ) from scimba_jax.nonlinear_approximation.networks.mlp import MLP from scimba_jax.nonlinear_approximation.numerical_solvers.projectors import Projector from scimba_jax.physical_models.elliptic_pde.grad_shafranov import ( GradShafranovFullPhysics2DWithXPoints, ) from scimba_jax.physical_models.elliptic_pde.grad_shafranov_fem import ( EQDSKBoundaryFlux, GradShafranovFullPhysicsNewtonWeakModel, GradShafranovFullPhysicsWeakModel, ) _HERE = Path(__file__).parent _DATA = _HERE.parents[5] / "src/scimba_jax/domains/tokamak/data" # Per machine: the PINN example's physics constants (Nicolas Pailliez's values), # the epochs, the FEM mesh size and the finer reference mesh. MACHINES = { "jet": { "eqdsk": _DATA / "eqdsk_jet_compare.dat", "k1": 0.1, "mu_fixed": 1.0, "epochs": (80, 600), "xpoint_down": (2.5, -1.5), "xpoint_up": (2.5, 2.0), "axis_x0": (3.1, 0.4), "mesh_size": 0.1, "reference_mesh_size": 0.05, }, "mast": { "eqdsk": _DATA / "g_p49320_t0.60000", "k1": 1.05, "mu_fixed": 1.2, "epochs": (200, 500), "xpoint_down": None, "xpoint_up": None, "axis_x0": (1.0, 0.0), "mesh_size": 0.1, "reference_mesh_size": 0.05, }, } K2 = 1.0 W_BC = 300.0 N_COLLOC = 5_000 N_BC_COLLOC = 5_000 HIDDEN = [12] * 4 ORDER = 2 # FEM degree and geometry degree PICARD_TOL = 1e-10 def _lagrange(y, i, mesh): return local_lagrange_basis(y, i, mesh, order=ORDER, out_dim=1) def _lagrange_by_logical(y, i, mesh): return local_lagrange_basis_by_logical(y, i, mesh, order=ORDER, out_dim=1) def wall_mesh(name, eqdsk, mesh_size): """The Q2 mesh of the wall, from a cache: GMSH does not repeat itself. Two ``macro_mesh_from_points`` calls on the same wall give the same cell count and different nodes (see ``laplacian/solve_laplacian_unstructured_2d .py``), so a comparison would change from one run to the next. Args: name: Machine name, for the cache file. eqdsk: The EQDSK file. mesh_size: Target cell size (m). Returns: The ``UnstructuredMesh``. """ cache = _HERE / f"mesh_{name}_p{ORDER}_h{mesh_size}.npz" if not cache.exists(): macro = macro_mesh_from_points( eqdsk_wall_loop(str(eqdsk)), order=ORDER, mesh_size=mesh_size, smooth=True ) np.savez_compressed(cache, nodes=macro.nodes, cells=macro.cells) stored = np.load(cache) return UnstructuredMesh( nodes=stored["nodes"], cells=stored["cells"], ref_quad=UnitSquareTensorized(dim=2, order=4), order=ORDER, ) def _fem_variables(machine, mesh_size): """Q2 Lagrange space on the cached wall mesh of size ``mesh_size``.""" mesh = wall_mesh(machine["name"], machine["eqdsk"], mesh_size) basis = AnalyticBasis( nb_basis=(ORDER + 1) ** 2, out_dim=1, mesh=mesh, local_basis=_lagrange, local_basis_by_logical=_lagrange_by_logical, basis_type="scalar", ) return VariablesFE(basis=basis, nb_variables=1, dof_map=UnstructuredLagrangeDofMap) def _fem_model( machine, eqdsk_data, wall, nonlinear, cls=GradShafranovFullPhysicsWeakModel ): """The FEM model with the machine's settings (Picard class by default).""" r_grid, z_grid, psi_2d, ff_prime, p_prime = eqdsk_data return cls( ff_prime, p_prime, psi_2d, r_grid, z_grid, k1=machine["k1"], k2=K2, nonlinear=nonlinear, R_wall=wall[0], Z_wall=wall[1], xpoint_down_target=machine["xpoint_down"], xpoint_up_target=machine["xpoint_up"], newton_axis_x0=machine["axis_x0"], mu_fixed=machine["mu_fixed"], ) NEWTON_SETTINGS = {"cg_solver": "bicgstab", "line_search": "armijo"} def solve_reference(machine, eqdsk_data, wall): """The same model on a finer mesh, to Newton accuracy: the reference. Linear phase, five Picard steps, then the exact Newton to ``||F|| < 1e-10``. Compared at the coarse nodes, it separates the discretisation error of the coarse FEM (and the PINN's error) from the EQDSK file's own. Returns: The solved variables (evaluable anywhere) and the Newton report. """ variables = _fem_variables(machine, machine["reference_mesh_size"]) begin = time.perf_counter() linear = EllipticFEscheme.solve( EllipticFEscheme(_fem_model(machine, eqdsk_data, wall, False), variables) ) picard = EllipticFEscheme(_fem_model(machine, eqdsk_data, wall, True), variables) warm = EllipticFEscheme.solve( picard, dofsl_init=jnp.array(linear.variables.dofsl), max_iter=5 ) newton = EllipticFEscheme( _fem_model( machine, eqdsk_data, wall, True, GradShafranovFullPhysicsNewtonWeakModel ), variables, linearization="jvp", ) solved, report = EllipticFEscheme.solve( newton, dofsl_init=jnp.array(warm.variables.dofsl), return_report=True, tol=PICARD_TOL, max_iter=40, **NEWTON_SETTINGS, ) jax.block_until_ready(solved.variables.dofsl) print( f" reference h = {machine['reference_mesh_size']}: " f"{variables.mesh.n_cells_total} cells, {variables.dofsl.shape[0]} DOFs, " f"Picard x5 + {int(report.n_iter)} Newton, ||F|| = " f"{float(report.residual):.1e}, {time.perf_counter() - begin:.1f} s" ) return solved.variables, report def solve_fem(machine, eqdsk_data, wall): """The linear phase, then the full physics three ways. * **Picard**: :class:`GradShafranovFullPhysicsWeakModel` (source and axis / X-points frozen), stiffness only in the linearisation, CG, blocks; * **Newton**: :class:`GradShafranovFullPhysicsNewtonWeakModel` (source and axis / X-point fluxes differentiated), matrix-free, BiCGSTAB (the Jacobian is not symmetric), Armijo, straight from the linear phase; * **Picard x5 + Newton**: the same Newton after five Picard steps. Each solve is run twice: the first time includes the compilation, the second is the run alone. Returns: ``(nodes, results, pre)``: node positions; per method a dict with the nodal flux, iterations, residual, convergence and both times; the Picard solution's axis and X-point quantities. """ variables = _fem_variables(machine, machine["mesh_size"]) mesh = variables.mesh def model(nonlinear, cls=GradShafranovFullPhysicsWeakModel): return _fem_model(machine, eqdsk_data, wall, nonlinear, cls) linear = EllipticFEscheme.solve(EllipticFEscheme(model(False), variables)) initial = jnp.array(linear.variables.dofsl) print(f" {mesh.n_cells_total} cells, {variables.dofsl.shape[0]} DOFs") picard = EllipticFEscheme(model(True), variables) newton = EllipticFEscheme( model(True, GradShafranovFullPhysicsNewtonWeakModel), variables, linearization="jvp", ) newton_settings = NEWTON_SETTINGS def timed(scheme, start, **settings): # ⚠ The solves share `variables` and write the solution in place: # everything is copied out before the next one. elapsed = [] for _ in range(2): # the first call compiles begin = time.perf_counter() solved, report = EllipticFEscheme.solve( scheme, dofsl_init=start, return_report=True, tol=PICARD_TOL, **settings ) # JAX returns before the work is done: wait for it. jax.block_until_ready(solved.variables.dofsl) elapsed.append(time.perf_counter() - begin) return { "psi": np.array(solved.variables.dofsl[:, 0]), "iterations": int(report.n_iter), "residual": float(report.residual), "converged": bool(report.converged), "time_compiled": elapsed[0], "time": elapsed[1], } results = {"Picard": timed(picard, initial, max_iter=200)} results["Newton"] = timed(newton, initial, max_iter=40, **newton_settings) warm = EllipticFEscheme.solve(picard, dofsl_init=initial, max_iter=5) start = jnp.array(warm.variables.dofsl) results["Picard x5 + Newton"] = timed(newton, start, max_iter=40, **newton_settings) results["Picard x5 + Newton"]["iterations"] += 5 for name, result in results.items(): print( f" {name:20s} {result['iterations']:4d} iterations, " f"||F|| = {result['residual']:.1e}" f"{'' if result['converged'] else ' (NOT converged)'}, " f"{result['time_compiled']:.2f} s with compilation, " f"{result['time']:.2f} s without" ) pre = jax.jit(picard.pde.pre_computation_without_diff)( variables, jnp.asarray(results["Picard"]["psi"])[:, None] ) nodes = np.asarray( variables.boundary_dof_positions(np.arange(variables.dofsl.shape[0])) ) return nodes, results, pre def solve_pinn(machine, eqdsk_data, wall): """The PINN example's two phases, with the settings above. Returns: ``(space, model, sampler, seconds)`` after phase 2. """ r_grid, z_grid, psi_2d, ff_prime, p_prime = eqdsk_data r_wall, z_wall = wall domain = Square2D( [ [float(r_wall.min()), float(r_wall.max())], [float(z_wall.min()), float(z_wall.max())], ], is_main_domain=True, ) sampler = TokamakSampler2D(r_wall, z_wall, oversample=5) key, subkey = jax.random.split(jax.random.PRNGKey(0)) network = MLP(in_size=2, out_size=1, hidden_sizes=HIDDEN, key=subkey) space = ApproximationSpace({"x": 2}, [(network, "scalar", None)], model_type="x") def model(nonlinear): return GradShafranovFullPhysics2DWithXPoints( domain, FFprime=ff_prime, Pprime=p_prime, psi_2d=psi_2d, R_grid=r_grid, Z_grid=z_grid, k1=machine["k1"], k2=K2, mu_fixed=machine["mu_fixed"], nonlinear=nonlinear, bc="weak", model_type="x", R_wall=r_wall, Z_wall=z_wall, xpoint_down_target=machine["xpoint_down"] if nonlinear else None, xpoint_up_target=machine["xpoint_up"] if nonlinear else None, newton_axis_x0=machine["axis_x0"] if nonlinear else None, ) epochs_1, epochs_2 = machine["epochs"] start = time.perf_counter() phase_1 = Projector( model(False), space, sampler, weights={"interior": [1.0], "boundary": [W_BC]} ) key, phase_1 = phase_1.project(key, space, epochs_1, N_COLLOC, N_BC_COLLOC) print(f" phase 1: loss {phase_1.best_loss['total']:.2e}") full = model(True) phase_2 = Projector( full, phase_1.space, sampler, weights={"interior": [1.0], "boundary": [W_BC]}, matrix_regularization=9.0e-6, ) key, phase_2 = phase_2.project(key, phase_1.space, epochs_2, N_COLLOC, N_BC_COLLOC) elapsed = time.perf_counter() - start print( f" phase 2: loss {phase_2.best_loss['total']:.2e}, " f"{elapsed:.0f} s for both phases (compilation included)" ) return phase_2.space, full, sampler, elapsed def _describe(label, pre): has = np.asarray([pre["has_down"], pre["has_up"]], dtype=bool) xpoints = np.asarray(pre["x_xpoint"])[has] print( f" {label:6s} psi_axis {float(pre['psi_axis']):+.4f} at " f"{np.round(np.asarray(pre['x_axis']), 3)}, psi_bnd " f"{float(pre['psi_bnd']):+.4f}, X-points {np.round(xpoints, 3).tolist()}" ) def _panel(ax, triangulation, field, title, wall, markers=(), cmap="turbo"): tc = ax.tricontourf(triangulation, field, levels=40, cmap=cmap) plt.colorbar(tc, ax=ax) ax.plot(*np.vstack([np.stack(wall, 1), np.stack(wall, 1)[:1]]).T, "k-", lw=1.0) for pre, color in markers: ax.plot(*np.asarray(pre["x_axis"]), "+", color=color, ms=12, mew=2) for k, has in enumerate((pre["has_down"], pre["has_up"])): if bool(has): ax.plot(*np.asarray(pre["x_xpoint"])[k], "x", color=color, ms=9, mew=2) # The wall's box: a marker found outside (an unconverged PINN) must not # rescale the panel. ax.set_xlim(np.min(wall[0]) - 0.05, np.max(wall[0]) + 0.05) ax.set_ylim(np.min(wall[1]) - 0.05, np.max(wall[1]) + 0.05) ax.set_aspect("equal") ax.set_title(title, fontsize=9) ax.set_xlabel("R [m]") ax.set_ylabel("Z [m]") def run(name): """Solve by FEM and by PINN, compare, plot. Args: name: ``"jet"`` or ``"mast"``. """ machine = dict(MACHINES[name], name=name) eqdsk_data = read_eqdsk(str(machine["eqdsk"]), normalize=False) wall = read_eqdsk_wall(str(machine["eqdsk"])) print(f"\n{name.upper()} — FEM (Q{ORDER})") nodes, fem, pre_fem = solve_fem(machine, eqdsk_data, wall) psi_fem = fem["Picard"]["psi"] reference, _ = solve_reference(machine, eqdsk_data, wall) psi_ref = np.asarray(reference.evaluate(jnp.asarray(nodes)))[:, 0] print(f"\n{name.upper()} — PINN") space, pinn_model, sampler, pinn_time = solve_pinn(machine, eqdsk_data, wall) points = jnp.asarray(nodes) psi_pinn = np.asarray( space.create_variables()[0].vmap_on_physical_variables()(space, points)[:, 0] ) _, sample = sampler.sample(jax.random.PRNGKey(999), N_COLLOC, N_BC_COLLOC) pre_pinn = pinn_model.pre_computation_without_diff(space, sample) r_grid, z_grid, psi_2d, _, _ = eqdsk_data psi_eqdsk = np.asarray(jax.vmap(EQDSKBoundaryFlux(psi_2d, r_grid, z_grid))(points))[ :, 0 ] scale = np.max(np.abs(psi_eqdsk)) def gap(a, b): return np.max(np.abs(a - b)) / scale print(f"\n{name.upper()} — comparison, max over the FEM nodes / max|psi_EQDSK|") print( f" {'':20s} {'iter.':>5s} {'compiled':>9s} {'run':>7s} {'vs Picard':>10s} " f"{'vs PINN':>9s} {'vs EQDSK':>9s} {'vs fine FEM':>12s}" ) for method, result in fem.items(): flag = "" if result["converged"] else " (not converged)" print( f" {method:20s} {result['iterations']:5d} " f"{result['time_compiled']:8.2f}s {result['time']:6.2f}s " f"{gap(result['psi'], psi_fem):10.1e} {gap(result['psi'], psi_pinn):9.1e} " f"{gap(result['psi'], psi_eqdsk):9.1e} " f"{gap(result['psi'], psi_ref):12.1e}{flag}" ) print( f" {'PINN':20s} {'':5s} {pinn_time:8.0f}s {'':7s} " f"{gap(psi_pinn, psi_fem):10.1e} " f"{'':9s} {gap(psi_pinn, psi_eqdsk):9.1e} {gap(psi_pinn, psi_ref):12.1e}" ) print( f" {'EQDSK':20s} {'':5s} {'':9s} {'':7s} {gap(psi_eqdsk, psi_fem):10.1e} " f"{gap(psi_eqdsk, psi_pinn):9.1e} {'':9s} {gap(psi_eqdsk, psi_ref):12.1e}" ) _describe("Picard", pre_fem) _describe("PINN", pre_pinn) # Delaunay on the nodes of a non-convex wall: mask what falls outside. triangulation = Triangulation(nodes[:, 0], nodes[:, 1]) centres = nodes[triangulation.triangles].mean(axis=1) triangulation.set_mask(~PolygonPath(np.stack(wall, 1)).contains_points(centres)) fig, axes = plt.subplots(2, 3, figsize=(15, 11), constrained_layout=True) markers = ((pre_fem, "white"), (pre_pinn, "magenta")) _panel( axes[0, 0], triangulation, psi_fem, "psi FEM Picard (+ axis, x X-points)", wall, markers[:1], ) _panel(axes[0, 1], triangulation, psi_pinn, "psi PINN", wall, markers[1:]) _panel(axes[0, 2], triangulation, psi_eqdsk, "psi EQDSK", wall) for ax, field, label, marks in ( (axes[1, 0], psi_fem, "FEM Picard", markers[:1]), (axes[1, 1], psi_pinn, "PINN", markers[1:]), (axes[1, 2], psi_eqdsk, "EQDSK", ()), ): _panel( ax, triangulation, np.abs(field - psi_ref), f"|{label} - fine FEM (h = {machine['reference_mesh_size']})|", wall, marks, cmap="viridis", ) fig.suptitle(f"Grad-Shafranov fixed boundary, {name.upper()}: FEM Q{ORDER} vs PINN") out = _HERE / f"grad_shafranov_{name}_fem_vs_pinn.png" fig.savefig(out, dpi=130) print(f"\nSaved figure to {out}") plt.show()