r"""Anisotropic diffusion in the JET tokamak, solved with a time-discrete PINN. Toroidal coordinates ``x = (R, Z, phi)``: .. math:: \partial_t \rho = (D_\parallel - D_\perp)\, \nabla_\parallel^2 \rho + D_\perp\, \nabla^2 \rho , with the field-aligned direction ``b`` reconstructed from the poloidal flux ``psi(R, Z)`` of the JET equilibrium (EQDSK). Strong anisotropy ``D_par >> D_perp`` makes an initial blob spread along the flux surfaces while barely crossing them. Two ways to obtain the base flux ``psi`` are provided via ``PSI_SOURCE``: - ``"eqdsk"`` (default, faithful to Nicolas Pailliez's ``diffaniso_jet.py``): fit a smooth network to the EQDSK ``psi_2D`` field (a regression, not a PDE solve), via :class:`FunctionApproximator`. - ``"grad_shafranov"``: solve the Grad-Shafranov PINN (:class:`GradShafranovFullPhysics2DWithXPoints`) to get a PDE-consistent psi. Time discretization: SDIRK2 (Pareschi-Russo, gamma = 1 - 1/sqrt(2)) via :class:`~scimba_jax.nonlinear_approximation.numerical_solvers.discrete_pinns.DiscretePINN` -- the same 2-stage L-stable scheme Nicolas uses in torch. Ported from Nicolas Pailliez's ``diffaniso_jet.py`` (``new_scimba_nicolas``, scimba-torch) to scimba-JAX. """ 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 MplPath from scimba_jax.domains.meshless_domains.domains_2d import Square2D from scimba_jax.domains.meshless_domains.domains_nd import HypercubeND from scimba_jax.domains.tokamak import ( TokamakSampler2D, TokamakSampler3DFrom2D, read_eqdsk, read_eqdsk_wall, ) 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.discrete_pinns import ( DiscretePINN, ) from scimba_jax.nonlinear_approximation.numerical_solvers.projectors import Projector from scimba_jax.physical_models.function_approximator.function_approximator import ( FunctionApproximator, ) from scimba_jax.physical_models.temporal_pde.diffaniso_tokamak import ( SteadyAnisotropicDiffusionTokamak, ) from scimba_jax.time_discrete.butcher_tableau import ( build_pareschi_russo_tableau, ) # ── configuration ───────────────────────────────────────────────────────────── PSI_SOURCE = "eqdsk" # "eqdsk" (regression, faithful) or "grad_shafranov" D_PAR = 1.0 # parallel diffusion (Nicolas: mu1) D_PERP = 0.0 # perpendicular diffusion (Nicolas: mu2) FINAL_TIME = 1e-0 DT_MARCH = 1e-2 NT = 100 N_COLLOC_PSI = 6000 N_COLLOC_INIT = 30000 # torch value: needed to resolve the narrow sigma=0.08 peak N_COLLOC_TIME = 30000 N_EPOCHS_PSI = 300 N_EPOCHS_INIT = 500 N_EPOCHS = 10 _DATA = Path(__file__).resolve().parents[4] / "src/scimba_jax/domains/tokamak/data" _EQDSK = _DATA / "eqdsk_jet_compare.dat" # ── read the JET equilibrium ────────────────────────────────────────────────── R_grid, Z_grid, psi_2D, _FFp, _Pp = read_eqdsk(str(_EQDSK), normalize=False) R_wall, Z_wall = read_eqdsk_wall(str(_EQDSK)) def _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() ] lines = Path(path).read_text().splitlines() r_ref = chunks(lines[1])[2] b_ref = chunks(lines[2])[4] return b_ref, r_ref _b_ref, _r_ref = _eqdsk_header(str(_EQDSK)) F0 = _b_ref * _r_ref # toroidal flux parameter: B_phi = F0 / R print(f"JET: F0 = {F0:.4f} wall R∈[{R_wall.min():.2f},{R_wall.max():.2f}]") R_MIN, R_MAX = float(R_wall.min()), float(R_wall.max()) Z_MIN, Z_MAX = float(Z_wall.min()), float(Z_wall.max()) key = jax.random.PRNGKey(0) # ── base flux psi ───────────────────────────────────────────────────────────── def _build_psi_eqdsk(key): """Fit a smooth network to the EQDSK psi_2D field (regression).""" rg = jnp.asarray(R_grid) zg = jnp.asarray(Z_grid) psi_grid = jnp.asarray(psi_2D) # shape (nZ, nR) def psi_target(x): # bilinear interpolation of psi_2D at (R, Z), differentiable via jax. fr = (x[0] - rg[0]) / (rg[-1] - rg[0]) * (rg.shape[0] - 1) fz = (x[1] - zg[0]) / (zg[-1] - zg[0]) * (zg.shape[0] - 1) val = jax.scipy.ndimage.map_coordinates( psi_grid, jnp.stack([fz, fr]), order=1, mode="nearest" ) return jnp.reshape(val, (1,)) dom2d = Square2D([[R_MIN, R_MAX], [Z_MIN, Z_MAX]], is_main_domain=True) sampler2d = TokamakSampler2D(R_wall, Z_wall, oversample=5) model = FunctionApproximator(dom2d, 1, "x", lambda x: psi_target(x)) k, sub = jax.random.split(key) nn = MLP(in_size=2, out_size=1, hidden_sizes=[20] * 4, key=sub) sp = ApproximationSpace({"x": 2}, [(nn, "scalar", None)], model_type="x") pinn = Projector(model, sp, sampler2d) k, pinn = pinn.project(k, sp, N_EPOCHS_PSI, N_COLLOC_PSI) print(f" psi regression best loss: {pinn.best_loss['total']:.3e}") return k, pinn.space def _build_psi_grad_shafranov(key): """Solve the Grad-Shafranov PINN to get a PDE-consistent psi (JET).""" from scimba_jax.physical_models.elliptic_pde.grad_shafranov import ( GradShafranovFullPhysics2DWithXPoints, ) dom2d = Square2D([[R_MIN, R_MAX], [Z_MIN, Z_MAX]], is_main_domain=True) sampler2d = TokamakSampler2D(R_wall, Z_wall, oversample=5) model = GradShafranovFullPhysics2DWithXPoints( dom2d, FFprime=_FFp, Pprime=_Pp, psi_2D=psi_2D, R_grid=R_grid, Z_grid=Z_grid, k1=0.1, k2=1.0, mu_fixed=1.0, nonlinear=False, bc="weak", model_type="x", R_wall=R_wall, Z_wall=Z_wall, ) k, sub = jax.random.split(key) nn = MLP(in_size=2, out_size=1, hidden_sizes=[16] * 4, key=sub) sp = ApproximationSpace({"x": 2}, [(nn, "scalar", None)], model_type="x") pinn = Projector( model, sp, sampler2d, weights={"interior": [1.0], "boundary": [100.0]} ) k, pinn = pinn.project(k, sp, N_EPOCHS_PSI, N_COLLOC_PSI, N_COLLOC_PSI) print(f" grad-shafranov best loss: {pinn.best_loss['total']:.3e}") return k, pinn.space print(f"\nBuilding base psi (source = {PSI_SOURCE}) ...") if PSI_SOURCE == "grad_shafranov": key, psi_space = _build_psi_grad_shafranov(key) else: key, psi_space = _build_psi_eqdsk(key) (_psi_var,) = psi_space.create_variables() def psi_func(x: jnp.ndarray) -> jnp.ndarray: """Frozen base flux psi(R, Z) from the trained network (ignores phi).""" return _psi_var(jax.lax.stop_gradient(psi_space), x[:2]) # ── anisotropic diffusion: domain, model, discrete PINN ─────────────────────── DOM_X = HypercubeND( [(R_MIN, R_MAX), (Z_MIN, Z_MAX), (-jnp.pi, jnp.pi)], is_main_domain=True ) SAMPLER = TokamakSampler3DFrom2D(R_wall, Z_wall, phi_range=(-jnp.pi, jnp.pi)) DOM_T = (0.0, FINAL_TIME) # Gaussian blob on a constant baseline, torch-faithful constants (Nicolas: # R0=3.335, Z0=0.25, phi0=0, sigma=0.08, a=20). The initial condition is # ``1 + a * exp(...)`` with a periodic wrap of phi. R0, Z0, PHI0, SIGMA, AMP = 3.335, 0.25, 0.0, 0.08, 20.0 def f_init(x: jnp.ndarray) -> jnp.ndarray: r, z, phi = x[0], x[1], x[2] dphi = jnp.arctan2(jnp.sin(phi - PHI0), jnp.cos(phi - PHI0)) return ( 1.0 + AMP * jnp.exp(-((r - R0) ** 2 + (R0 * dphi) ** 2 + (z - Z0) ** 2) / SIGMA**2) ).reshape(1) pde = SteadyAnisotropicDiffusionTokamak( main_domain=DOM_X, time_domain=DOM_T, psi_func=psi_func, f0=F0, d_par=D_PAR, d_perp=D_PERP, bc="none", # no wall BC: D_perp~0 -> pure parallel transport, phi-periodic ) discrete_pinn = DiscretePINN( DOM_X, DOM_T, SAMPLER, 1, NT, build_pareschi_russo_tableau(), None, pde ) key, subkey = jax.random.split(key) nn = MLP(in_size=8, out_size=1, hidden_sizes=[20] * 5, key=subkey) def preprocessing_torus(x: jnp.ndarray) -> jnp.ndarray: """Toroidal preprocessing: Fourier features of phi (modes 1, 2, 3).""" R, Z, phi = x[0], x[1], x[2] return jnp.stack( [ R, Z, jnp.sin(phi), jnp.cos(phi), jnp.sin(2.0 * phi), jnp.cos(2.0 * phi), jnp.sin(3.0 * phi), jnp.cos(3.0 * phi), ] ) space = ApproximationSpace( {"x": 3}, [(nn, "scalar", None)], model_type="x", pre_processing=preprocessing_torus ) print("\nFitting the initial condition ...") key, discrete_pinn = discrete_pinn.initialize( key, space, f_init, N_EPOCHS_INIT, N_COLLOC_INIT, matrix_regularization=1e-8 ) space = discrete_pinn.space print(f"Time-marching to t = {FINAL_TIME} ...") key, discrete_pinn = discrete_pinn.solve(key, space, N_EPOCHS, N_COLLOC_TIME) new_space = discrete_pinn.space # ── plots: poloidal slices at t=0 (IC) and t=final, shared colour scale ─────── (u_fn,) = new_space.create_variables() (u_fn_init,) = space.create_variables() n_plot = 90 rs = np.linspace(R_MIN, R_MAX, n_plot) zs = np.linspace(Z_MIN, Z_MAX, n_plot) rr, zz = np.meshgrid(rs, zs) wall_path = MplPath(np.stack([R_wall, Z_wall], axis=1)) inside = wall_path.contains_points(np.stack([rr.ravel(), zz.ravel()], axis=1)) TOKAMAK = "JET" # psi background (flux surfaces = level lines of the equilibrium) for overlay psi_grid = np.array( jax.vmap(lambda p: psi_func(jnp.array([p[0], p[1], 0.0])))( jnp.stack([rr.ravel(), zz.ravel()], axis=-1) ) ).reshape(n_plot, n_plot) psi_grid[~inside.reshape(n_plot, n_plot)] = np.nan # Toroidal slices starting at phi=0 (where the blob sits): 0, 40, ..., 320 deg. phi_slices = np.linspace(0.0, 2.0 * np.pi, 10)[:-1] def _eval_slices(space_, ufn_): """Evaluate rho on the (rr, zz) grid for every phi slice (NaN outside wall).""" grids = [] for phi in phi_slices: pts = jnp.stack([rr.ravel(), zz.ravel(), jnp.full(rr.size, phi)], axis=-1) u = np.array(jax.vmap(lambda p: ufn_(space_, p))(pts)).reshape(-1) u[~inside] = np.nan grids.append(u.reshape(n_plot, n_plot)) return grids u_grids_init = _eval_slices(space, u_fn_init) # t = 0 (initial condition) u_grids = _eval_slices(new_space, u_fn) # t = final (also used by the 3-D render) # colour scale for the 3-D render (final-time solution) vmin = float(np.nanmin(u_grids)) vmax = float(np.nanmax(u_grids)) # quantitative check: does the blob peak move between t=0 and t=final? c0 = float(jnp.squeeze(u_fn_init(space, jnp.array([R0, Z0, 0.0])))) cT = float(jnp.squeeze(u_fn(new_space, jnp.array([R0, Z0, 0.0])))) print(f"rho(blob center): t=0 -> {c0:.4f} t={FINAL_TIME} -> {cT:.4f}") def _plot_poloidal(grids, t_label, tag): # each slice gets its own adaptive vmin/vmax and its own colorbar (not # shared across the 9 panels) -- every phi angle is scaled to its own # local min/max so its structure is visible regardless of the other slices. fig, axes = plt.subplots(3, 3, figsize=(15, 12), constrained_layout=True) for ax, phi, u_grid in zip(axes.ravel(), phi_slices, grids): cf = ax.pcolormesh(rr, zz, u_grid, shading="auto", cmap="turbo") # only psi flux surfaces (the Gaussian is too narrow to contour) ax.contour( rr, zz, psi_grid, levels=10, colors="0.5", linewidths=0.35, alpha=0.6 ) ax.plot(list(R_wall) + [R_wall[0]], list(Z_wall) + [Z_wall[0]], "k-", lw=1.2) ax.set_title(rf"$\varphi = {np.degrees(phi):.0f}\degree$", fontsize=9) ax.set_xlabel("R") ax.set_ylabel("Z") ax.set_aspect("equal") fig.colorbar(cf, ax=ax, shrink=0.85, label="rho") plt.suptitle( f"{TOKAMAK} anisotropic diffusion — poloidal slices at t={t_label} " f"(psi: {PSI_SOURCE})", fontsize=13, ) plt.savefig( f"diffaniso_{TOKAMAK.lower()}_poloidal_{tag}.png", dpi=130, bbox_inches="tight" ) _plot_poloidal(u_grids_init, 0.0, "t0") _plot_poloidal(u_grids, FINAL_TIME, "final") # ── toroidal-average plots (Nicolas): _phi - baseline, psi contours only ── # The blob is a localized spot at t=0; its toroidal average spreads into a ring # along the flux surface as it diffuses. Only the psi level lines are drawn (the # Gaussian itself is too narrow to contour meaningfully). N_PHI_AVG = 48 phi_avg = np.linspace(0.0, 2.0 * np.pi, N_PHI_AVG, endpoint=False) def _toroidal_mean(space_, ufn_): """Toroidal average of rho on the (rr, zz) grid, minus the baseline (1).""" acc = np.zeros(rr.size) for phi in phi_avg: pts = jnp.stack([rr.ravel(), zz.ravel(), jnp.full(rr.size, phi)], axis=-1) acc += np.array(jax.vmap(lambda p: ufn_(space_, p))(pts)).reshape(-1) u = acc / N_PHI_AVG - 1.0 # toroidal mean minus baseline u[~inside] = np.nan return u.reshape(n_plot, n_plot) def _plot_toroidal_mean(u_tor, t_label, tag): fig, ax = plt.subplots(figsize=(6, 8), constrained_layout=True) cf = ax.pcolormesh(rr, zz, u_tor, shading="auto", cmap="turbo") # own scale # only the psi flux surfaces (the blob is too narrow to contour) ax.contour(rr, zz, psi_grid, levels=12, colors="0.5", linewidths=0.4, alpha=0.7) ax.plot(list(R_wall) + [R_wall[0]], list(Z_wall) + [Z_wall[0]], "k-", lw=1.2) ax.set_title( rf"{TOKAMAK} toroidal mean $\langle\rho\rangle_\varphi - 1$ (t={t_label})", fontsize=11, ) ax.set_xlabel("R") ax.set_ylabel("Z") ax.set_aspect("equal") fig.colorbar(cf, ax=ax, shrink=0.85, label=r"$\langle\rho\rangle_\varphi - 1$") plt.savefig( f"diffaniso_{TOKAMAK.lower()}_tormean_{tag}.png", dpi=130, bbox_inches="tight" ) _plot_toroidal_mean(_toroidal_mean(space, u_fn_init), 0.0, "t0") _plot_toroidal_mean(_toroidal_mean(new_space, u_fn), FINAL_TIME, "final") # ── 3-D toroidal render: filled poloidal cross-sections colored by rho, swept # along the torus, inside a transparent wall envelope (paper-style). ───────── from mpl_toolkits.mplot3d.art3d import Poly3DCollection # noqa: E402 # reuse the masked poloidal grid from the slices above for the cross-sections _ins2d = inside.reshape(n_plot, n_plot) def _u_on_slice(phi): """Evaluate rho on the (rr, zz) grid at toroidal angle phi (NaN outside wall).""" pts = jnp.stack([rr.ravel(), zz.ravel(), jnp.full(rr.size, phi)], axis=-1) u = np.array(jax.vmap(lambda p: u_fn(new_space, p))(pts)).reshape(n_plot, n_plot) u[~_ins2d] = np.nan return u def render_torus_3d(full: bool): """Render cross-section disks + transparent wall envelope over the torus. Args: full: if True sweep the whole torus (0..2pi), else a partial arc that lets one see inside the tube (paper-style "telescope" view). """ if full: phi_env = np.linspace(0.0, 2.0 * np.pi, 160) phi_cuts = np.linspace(0.0, 2.0 * np.pi, 9)[:-1] tag, azim = "full", -60 else: phi_env = np.linspace(0.1 * np.pi, 1.2 * np.pi, 120) phi_cuts = np.linspace(0.15 * np.pi, 1.15 * np.pi, 5) tag, azim = "arc", -70 norm = plt.Normalize(vmin, vmax) fig = plt.figure(figsize=(11, 9)) ax = fig.add_subplot(111, projection="3d") # transparent wall envelope: the D-shaped wall polygon revolved over phi_env Xe = R_wall[None, :] * np.cos(phi_env[:, None]) Ye = R_wall[None, :] * np.sin(phi_env[:, None]) Ze = np.broadcast_to(Z_wall[None, :], Xe.shape) ax.plot_surface( Xe, Ye, Ze, color="steelblue", alpha=0.10, rstride=2, cstride=3, linewidth=0, shade=False, ) # filled cross-sections colored by rho for phi in phi_cuts: u_grid = _u_on_slice(phi) Xc = rr * np.cos(phi) Yc = rr * np.sin(phi) Zc = zz verts, cols = [], [] for a in range(n_plot - 1): for b in range(n_plot - 1): blk = u_grid[a : a + 2, b : b + 2] if not np.isnan(blk).any(): verts.append( [ (Xc[a, b], Yc[a, b], Zc[a, b]), (Xc[a + 1, b], Yc[a + 1, b], Zc[a + 1, b]), (Xc[a + 1, b + 1], Yc[a + 1, b + 1], Zc[a + 1, b + 1]), (Xc[a, b + 1], Yc[a, b + 1], Zc[a, b + 1]), ] ) cols.append(plt.cm.turbo(norm(blk.mean()))) ax.add_collection3d(Poly3DCollection(verts, facecolors=cols, edgecolors="none")) # wall outline of this slice xw = np.append(R_wall, R_wall[0]) * np.cos(phi) yw = np.append(R_wall, R_wall[0]) * np.sin(phi) zw = np.append(Z_wall, Z_wall[0]) ax.plot(xw, yw, zw, color="0.25", lw=0.7, alpha=0.8) sm = plt.cm.ScalarMappable(cmap="turbo", norm=norm) sm.set_array([]) fig.colorbar(sm, ax=ax, shrink=0.55, pad=0.02, label="rho") ax.set_title( f"{TOKAMAK} anisotropic diffusion in the torus (t={FINAL_TIME}) — " f"{'full torus' if full else 'partial arc'}", fontsize=12, ) ax.set_xlabel("x") ax.set_ylabel("y") ax.set_zlabel("Z") ax.set_box_aspect((1, 1, 0.5)) ax.view_init(elev=22, azim=azim) plt.tight_layout() out = f"diffaniso_{TOKAMAK.lower()}_torus3d_{tag}.png" plt.savefig(out, dpi=130, bbox_inches="tight") return out out_arc = render_torus_3d(full=False) out_full = render_torus_3d(full=True) plt.show() print( f"Saved: diffaniso_{TOKAMAK.lower()}_poloidal_t0.png, " f"diffaniso_{TOKAMAK.lower()}_poloidal_final.png, " f"diffaniso_{TOKAMAK.lower()}_tormean_t0.png, " f"diffaniso_{TOKAMAK.lower()}_tormean_final.png, {out_arc}, {out_full}" )