r"""Grad-Shafranov PINN on a real tokamak cross-section (JET / MAST) with X-points. Equation: .. math:: \Delta^* \psi = -(k_1^2 R^2 + k_2^2)(1 + \psi_n) \quad \text{in } \Omega \psi = 0 \quad \text{on } \partial\Omega where :math:`\psi_n` is the X-point masked normalised poloidal flux: .. math:: \psi_n = \frac{\psi - \psi_\mathrm{axis}}{\psi_\mathrm{bnd} - \psi_\mathrm{axis}} with reflection across the separatrix level in the scrape-off layer (outside the last closed flux surface). Domain: tokamak poloidal cross-section defined by the wall contour read from an EQDSK equilibrium file. Interior points are drawn by rejection sampling inside the wall polygon (:class:`~scimba_jax.domains.tokamak.TokamakSampler2D`). Training strategy (two phases): 1. **Adam** on the linearised problem (ψ_n = 0, constant source) to get a first approximation. 2. **Adam** with the full nonlinear source where ψ_axis and the X-point quantities are recomputed each step via :meth:`~GradShafranov2DWithXPoints.pre_computation_without_diff` and frozen (``stop_gradient``). Usage:: python grad_shafranov_with_xpoints.py # uses JET EQDSK python grad_shafranov_with_xpoints.py mast # uses MAST-4 EQDSK python grad_shafranov_with_xpoints.py mast-u # uses MAST-U EQDSK Ported from Nicolas Pailliez's ``new_FBGS_jet.py`` / ``new_FBGS_mast.py`` (``new_scimba_nicolas`` branch, scimba-torch) to scimba-JAX. """ import sys import timeit from pathlib import Path import jax import matplotlib.pyplot as plt from scimba_jax.domains.meshless_domains.domains_2d import Square2D from scimba_jax.domains.tokamak import TokamakSampler2D, read_eqdsk, read_eqdsk_wall from scimba_jax.nonlinear_approximation.approximation_spaces.approximation_spaces import ( 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, ) # ── Select EQDSK file ────────────────────────────────────────────────────────── _DATA = ( Path(__file__).parent.parent.parent.parent.parent.parent / "src/scimba_jax/domains/tokamak/data" ) _EQDSK_MAP = { "jet": _DATA / "eqdsk_jet_compare.dat", "mast": _DATA / "g_p49320_t0.60000", "mast-u": _DATA / "g_p44849_t0.20000", } _tok = sys.argv[1].lower() if len(sys.argv) > 1 else "jet" _eqdsk_path = _EQDSK_MAP.get(_tok, _EQDSK_MAP["jet"]) print(f"\nUsing EQDSK: {_eqdsk_path.name} (tokamak: {_tok.upper()})") import numpy as np # noqa: E402 R_grid, Z_grid, psi_2D, FFprime, Pprime = read_eqdsk(str(_eqdsk_path), normalize=False) R_wall, Z_wall = read_eqdsk_wall(str(_eqdsk_path)) print( f"Wall: {len(R_wall)} pts " f"R∈[{R_wall.min():.2f}, {R_wall.max():.2f}] " f"Z∈[{Z_wall.min():.2f}, {Z_wall.max():.2f}]" ) print( f"EQDSK profiles: n={len(FFprime)} FF'∈[{FFprime.min():.2f},{FFprime.max():.2f}]" ) # ── X-point targets depuis le champ psi_2D EQDSK ────────────────────────────── # Trouver où |∇ψ| est minimal près de ψ = ψ_bnd (séparatrice) → X-point def _find_xpoint_target_from_eqdsk( psi_2D, R_grid, Z_grid, psi_bnd, Z_axis, below: bool, tol: float = 0.05 ): """Extraire la position approx du X-point depuis le champ EQDSK.""" dZ = np.gradient(psi_2D, Z_grid, axis=0) dR = np.gradient(psi_2D, R_grid, axis=1) grad_norm = np.sqrt(dR**2 + dZ**2) span = abs(psi_bnd - psi_2D.min()) near_bnd = np.abs(psi_2D - psi_bnd) < tol * span # Restreindre au bon côté de l'axe Z2D = Z_grid[:, None] * np.ones_like(psi_2D) if below: near_bnd &= Z2D < Z_axis else: near_bnd &= Z2D > Z_axis if not near_bnd.any(): return None grad_masked = np.where(near_bnd, grad_norm, np.inf) iz, ir = np.unravel_index(np.argmin(grad_masked), grad_masked.shape) return float(R_grid[ir]), float(Z_grid[iz]) # Lire psi_axis, Z_axis depuis l'en-tête EQDSK (ligne 3) def _read_eqdsk_header(path): def chunks(line): return [ float(line[i : i + 16]) for i in range(0, len(line), 16) if line[i : i + 16].strip() ] with open(path) as f: lines = list(f) _, _, _psi_axis, _psi_bnd, _ = chunks(lines[2]) _R_axis, _Z_axis, *_ = chunks(lines[2]) return _R_axis, _Z_axis, _psi_axis, _psi_bnd _R_ax_eqdsk, _Z_ax_eqdsk, _psi_ax_eqdsk, _psi_bnd_eqdsk = _read_eqdsk_header( str(_eqdsk_path) ) xpoint_down_target = _find_xpoint_target_from_eqdsk( psi_2D, R_grid, Z_grid, _psi_bnd_eqdsk, _Z_ax_eqdsk, below=True ) xpoint_up_target = _find_xpoint_target_from_eqdsk( psi_2D, R_grid, Z_grid, _psi_bnd_eqdsk, _Z_ax_eqdsk, below=False ) newton_axis_x0 = (_R_ax_eqdsk, _Z_ax_eqdsk) print(f"EQDSK axis : R={_R_ax_eqdsk:.3f} Z={_Z_ax_eqdsk:.3f}") print(f"X-point ↓ : {xpoint_down_target}") print(f"X-point ↑ : {xpoint_up_target}") # ── Parameters — par tokamak (valeurs de Nicolas) ────────────────────────────── # # JET (new_FBGS_jet.py) : k1=-0.1→0.01R²+1 mu_fixed=1.0 x0=(3.1,0.4) # epochs 1000 + 500 [20]*5 # MAST (new_FBGS_mast.py): k1=1.05→1.1R²+1 mu_fixed=1.2 x0=(1.0,0.0) # epochs 2000 + 1000 [20]*5 # _JET = _tok == "jet" K1 = 0.1 if _JET else 1.05 # linéaire source K2 = 1.0 MU_FIXED = 1.0 if _JET else 1.2 # scale pression dans RHS W_BC = 100.0 N_COLLOC = 8_000 N_BC_COLLOC = 6_000 N_EPOCHS_1 = 80 if _JET else 200 N_EPOCHS_2 = 600 if _JET else 500 HIDDEN = [16] * 4 _XPT_DOWN = (2.5, -1.5) if _JET else None _XPT_UP = (2.5, 2.0) if _JET else None _NEWTON_X0 = (3.1, 0.4) if _JET else (1.0, 0.0) # ── Domain & sampler ─────────────────────────────────────────────────────────── # We use a Disk2D as a placeholder domain for the AbstractPhysicalModel # (provides label "interior" and boundary "boundary"). # Actual sampling is done by TokamakSampler2D which rejects points # outside the wall polygon → matches the same dict keys. 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) # ── Network & approximation space ────────────────────────────────────────────── key = jax.random.PRNGKey(0) key, subkey = jax.random.split(key) nn = MLP(in_size=2, out_size=1, hidden_sizes=HIDDEN, key=subkey) space = ApproximationSpace({"x": 2}, [(nn, "scalar", None)], model_type="x") # ── Phase 1: linear source (ψ_n = 0) ────────────────────────────────────────── print("\n── Phase 1: Adam, source constante (ψ_n = 0) ──────────────────") model_phase1 = GradShafranovFullPhysics2DWithXPoints( domain, FFprime=FFprime, Pprime=Pprime, psi_2d=psi_2D, R_grid=R_grid, Z_grid=Z_grid, k1=K1, k2=K2, mu_fixed=MU_FIXED, nonlinear=False, # source constante k1²R²+k2², BC = ψ_eqdsk bc="weak", model_type="x", R_wall=R_wall, Z_wall=Z_wall, ) pinn1 = Projector( model_phase1, space, sampler, weights={ "interior": [1.0], "boundary": [W_BC], }, # Nicolas: lr=2e-2, betas=(0.9, 0.999) ) key, sample_dict = sampler.sample(key, N_COLLOC, N_BC_COLLOC) print(f" initial loss: {pinn1.evaluate_loss(space, sample_dict):.4e}") t0 = timeit.default_timer() key, pinn1 = pinn1.project(key, space, N_EPOCHS_1, N_COLLOC, N_BC_COLLOC) print( f" best loss: {pinn1.best_loss['total']:.4e}" f" | {timeit.default_timer() - t0:.1f}s" ) space_phase1 = pinn1.space # ── Phase 2: nonlinear source with X-point detection ────────────────────────── print("\n── Détection des X-points (hors JIT) ──────────────────────────────") model_phase2 = GradShafranovFullPhysics2DWithXPoints( domain, FFprime=FFprime, Pprime=Pprime, psi_2d=psi_2D, R_grid=R_grid, Z_grid=Z_grid, k1=K1, k2=K2, nonlinear=True, bc="weak", model_type="x", R_wall=R_wall, Z_wall=Z_wall, xpoint_down_target=_XPT_DOWN, xpoint_up_target=_XPT_UP, newton_axis_x0=_NEWTON_X0, mu_fixed=MU_FIXED, ) print("\n── Phase 2: ENG, RHS full-physics FF' + μ₀P' (ψ_N interpolé) ──────") pinn2 = Projector( model_phase2, space_phase1, sampler, weights={"interior": [1.0], "boundary": [W_BC]}, matrix_regularization=9.0e-6, # Nicolas: 1e-7 → 1e-8 (plus stable avec X-points) ) key, sample_dict = sampler.sample(key, N_COLLOC, N_BC_COLLOC) print(f" initial loss: {pinn2.evaluate_loss(space_phase1, sample_dict):.4e}") t0 = timeit.default_timer() key, pinn2 = pinn2.project(key, space_phase1, N_EPOCHS_2, N_COLLOC, N_BC_COLLOC) print( f" best loss: {pinn2.best_loss['total']:.4e}" f" | {timeit.default_timer() - t0:.1f}s" ) # ── Plots (style Nicolas : scatter + griddata + contours) ───────────────────── from scipy.interpolate import griddata # noqa: E402 def plot_psi_with_diagnostics( ax, space, model, sampler, R_wall, Z_wall, title, psi_min=None, psi_max=None, n_samples=50_000, n_grid=200, ): """Port de plot_psi_with_diagnostics de Nicolas (new_FBGS_jet/mast.py).""" import numpy as _np from matplotlib.path import Path as MplPath key_p = jax.random.PRNGKey(999) key_p, sd = sampler.sample(key_p, n_samples, 0) pts = sd["interior"][0] variables = space.create_variables() psi_vals = _np.array(variables[0].vmap_on_physical_variables()(space, pts)[:, 0]) R = _np.array(pts[:, 0]) Z = _np.array(pts[:, 1]) if psi_min is None: psi_min = psi_vals.min() if psi_max is None: psi_max = psi_vals.max() # Grille régulière pour contours Rg = _np.linspace(R.min(), R.max(), n_grid) Zg = _np.linspace(Z.min(), Z.max(), n_grid) RR, ZZ = _np.meshgrid(Rg, Zg) psi_grid = griddata((R, Z), psi_vals, (RR, ZZ), method="linear") wall_path = MplPath(_np.stack([R_wall, Z_wall], axis=1)) outside = ~wall_path.contains_points(_np.stack([RR.ravel(), ZZ.ravel()], axis=1)) psi_grid.ravel()[outside] = _np.nan # Scatter sc = ax.scatter( R, Z, c=psi_vals, cmap="turbo", s=10, alpha=0.6, vmin=psi_min, vmax=psi_max ) # Contours avec labels (Nicolas: plt.clabel) CS = ax.contour(RR, ZZ, psi_grid, levels=20, colors="k", linewidths=0.6, alpha=0.7) ax.clabel(CS, inline=True, fontsize=7) plt.colorbar(sc, ax=ax, label="psi") # Mur Rw = list(R_wall) + [R_wall[0]] Zw = list(Z_wall) + [Z_wall[0]] ax.plot(Rw, Zw, "k-", lw=1.5) # Diagnostics (axe, psi_bnd, X-points) — Nicolas: compute_diagnostics psi_ax_val = None if model is not None and model._nonlinear: pre = model.pre_computation_without_diff(space, sd) xa = _np.array(pre["x_axis"]) psi_ax_val = float(pre["psi_axis"]) psi_bnd_val = float(pre["psi_bnd"]) xpts = _np.array(pre["x_xpoint"]) # (2, 2) psixpt = _np.array(pre["psi_xpoint"]) # (2,) has_d = bool(pre["has_down"]) has_u = bool(pre["has_up"]) # Axe magnétique (Nicolas: grey +) ax.scatter( xa[0], xa[1], s=150, marker="+", color="grey", zorder=6, label="psi-axis" ) # psi_bnd point : X-point le plus proche de psi_axis (Nicolas: green *) d_d = abs(psixpt[0] - psi_ax_val) if has_d else _np.inf d_u = abs(psixpt[1] - psi_ax_val) if has_u else _np.inf xb = xpts[0] if (has_d and d_d <= d_u) else xpts[1] if has_u else None if xb is not None: ax.scatter( xb[0], xb[1], s=120, marker="*", color="green", zorder=6, label="psi_bnd", ) # X-points (Nicolas: orange + pour down, cyan + pour up) if has_d: ax.scatter( xpts[0, 0], xpts[0, 1], s=100, marker="+", color="orange", zorder=6, label="X-point down", ) if has_u: ax.scatter( xpts[1, 0], xpts[1, 1], s=100, marker="+", color="cyan", zorder=6, label="X-point up", ) ax.legend(fontsize=7, loc="upper right") # Console print(f"\n[{title}]") print(f" ψ_axis = {psi_ax_val:.4e} @ R={xa[0]:.4f} Z={xa[1]:.4f}") print(f" ψ_bnd = {psi_bnd_val:.4f}") if has_d: print( f" X-point ↓ : R={xpts[0, 0]:.3f} Z={xpts[0, 1]:.3f} ψ={psixpt[0]:.4f}" ) if has_u: print( f" X-point ↑ : R={xpts[1, 0]:.3f} Z={xpts[1, 1]:.3f} ψ={psixpt[1]:.4f}" ) # Titre style Nicolas t = f"{title}" if psi_ax_val is not None: t += f"\n$\\psi_{{axis}}$ = {psi_ax_val:.4e} | $R_{{axis}}$ = {xa[0]:.4f}, $Z_{{axis}}$ = {xa[1]:.4f}" ax.set_title(t, fontsize=9) ax.set_xlabel("R") ax.set_ylabel("Z") ax.set_xlim(float(R_wall.min()) - 0.1, float(R_wall.max()) + 0.1) ax.set_ylim(float(Z_wall.min()) - 0.1, float(Z_wall.max()) + 0.1) ax.set_aspect("equal", adjustable="box") fig, axes = plt.subplots(2, 2, figsize=(14, 10), constrained_layout=True) fig.suptitle(f"Grad-Shafranov PINN — {_tok.upper()} (full physics)", fontsize=13) # Pertes for ax, pinn, label in zip( axes[:, 0], [pinn1, pinn2], ["Phase 1 (linéaire k₁²R²+k₂²)", "Phase 2 (FF′ + μ₀P′)"], ): for k, v in pinn.losses.losses_history.items(): ax.semilogy(v, label=k) ax.set_title(f"Loss — {label}") ax.set_xlabel("epoch") ax.legend(fontsize=8) ax.grid(True, which="both", alpha=0.3) # Approximation ψ plot_psi_with_diagnostics( axes[0, 1], pinn1.space, None, sampler, R_wall, Z_wall, "ψ — Phase 1 (linéaire)", ) plot_psi_with_diagnostics( axes[1, 1], pinn2.space, model_phase2, sampler, R_wall, Z_wall, "ψ — Phase 2 (full physics)", ) out_name = f"grad_shafranov_xpoints_{_tok}.png" plt.savefig(out_name, dpi=150, bbox_inches="tight") plt.show() print(f"Saved: {out_name}")