"""Projecting a field onto the tokamak built in ``tokamak_mesh.py``. The first thing worth running on a new mesh, and the cheapest: a DG mass matrix is block diagonal, so the projection is a small solve per cell with nothing coupling them. It needs no faces and no flux -- what it does exercise is the cell map and its Jacobian at every quadrature point, which on a product mesh carried through a global map is exactly the part that could be wrong. The field is a perturbation localised on the outboard midplane at ``phi = 0``, and the pictures are the ones a plasma physicist would ask for: the value on the cut surface, then poloidal sections at increasing angle. ⚠ **Each stage is finished before the next begins.** JAX is asynchronous: an unblocked ``project`` returns a pending array in microseconds and the work actually happens later, wherever its values are first read -- which was inside the plotting, so the figure appeared to hang while the projection it was drawing had not run. ``block_until_ready`` is what makes a printed timing mean anything. """ import time import jax import jax.numpy as jnp import matplotlib.pyplot as plt import numpy as np import tokamak_mesh as geometry 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, basis_values, ) from scimba_jax.linear_approximation.meshes.unstructured_mesh import ( unit_cell_quad_points, ) from scimba_jax.linear_approximation.variables.variables_dg import VariablesDG BASIS_ORDER = 2 # A perturbation localised on the outboard midplane: Gaussian across the # poloidal plane, Gaussian along the angle. # # ⚠ The widths are set by what the mesh can RESOLVE, and the example PRINTS # whether they are: a projection cannot exceed the maximum of what it projects, # so an overshoot above 1 measures under-resolution directly. An early version # used a ball of radius 0.35 in SPACE, which spans 0.12 rad -- a third of one # toroidal cell of the coarse mesh -- and overshot by 20% (1.199) while the # neighbouring slices read exactly zero. That picture was about the sampling, # not about the mesh. Poloidal cells are `1 / SECTION_CELLS` across and toroidal # ones subtend `2 pi / N_TOROIDAL`, so both sigmas below are one to two cells. BLOB_OFFSET = 0.45 SIGMA_POLOIDAL = 0.25 SIGMA_TOROIDAL = 0.5 def make_blob(centre, sigma_poloidal=SIGMA_POLOIDAL, sigma_toroidal=SIGMA_TOROIDAL): """The perturbation, as a function of PHYSICAL points. Written in the torus's own coordinates -- poloidal distance from ``centre = (R, Z)`` and toroidal angle -- because that is where its widths mean something, but read off the physical point, so the projection sees the geometry the mesh actually carries rather than a parameter-space caricature. ⚠ Continuous across ``phi = 0``, which the periodic mesh requires: the angle enters as ``arctan2`` on ``(-pi, pi]`` and only through its square, so the two sides of the seam agree. """ centre_radius, centre_height = centre def field(points): radius = jnp.sqrt(points[..., 0] ** 2 + points[..., 1] ** 2) poloidal = (radius - centre_radius) ** 2 + (points[..., 2] - centre_height) ** 2 angle = jnp.arctan2(points[..., 1], points[..., 0]) return jnp.exp( -poloidal / (2.0 * sigma_poloidal**2) - angle**2 / (2.0 * sigma_toroidal**2) )[..., None] return field def make_basis(mesh, order=BASIS_ORDER): """Lagrange of degree ``order`` on the cell, in both of its forms.""" return AnalyticBasis( nb_basis=(order + 1) ** mesh.dim, out_dim=1, mesh=mesh, basis_type="scalar", local_basis=lambda y, i, m: local_lagrange_basis( y, i, m, order=order, out_dim=1 ), # The by-logical form evaluates from a unit-cell preimage, with no # inverse map -- the fast path on a curved cell, and the only one the # pictures below use. local_basis_by_logical=lambda y, i, m: local_lagrange_basis_by_logical( y, i, m, order=order, out_dim=1 ), ) def evaluate(basis, variables, cells, reference): """The expansion on a batch of cells, at given UNIT-CELL points. ⚠ Passing ``x_hat`` is not an optimisation, it is the difference between working and not: a curved cell has no closed-form inverse map, so asking the basis about a bare physical point sends it into a Newton solve per point. Every point here comes from a reference grid, so its preimage is free. Args: basis: The basis the DOFs belong to. variables: Holder of ``dofsl``. cells: ``(n,)`` cell indices. reference: ``(n, q, dim)`` unit-cell points, per cell. Returns: ``(physical, values)``, shapes ``(n, q, dim)`` and ``(n, q)``. """ mesh = basis.mesh def per_cell(cell, unit_points): physical = mesh._unit_hypercube_to_cell(cell, unit_points) values = jax.vmap(lambda x_hat, x: basis_values(basis, cell, x, x_hat)[0])( unit_points, physical ) return physical, jnp.einsum("iv,qiv->q", variables.dofsl[cell], values) physical, values = jax.vmap(per_cell)(cells, reference) return np.asarray(physical), np.asarray(values) def projection_error(mesh, basis, variables, field): """Relative ``L2`` error, on the mesh's own quadrature.""" x_hat = unit_cell_quad_points(mesh) cells = geometry.all_cells(mesh) reference = jnp.broadcast_to(x_hat, (cells.shape[0],) + x_hat.shape) points, got = evaluate(basis, variables, cells, reference) weights = jax.vmap(lambda c: mesh._local_weights_points(c)[0])(cells) want = field(jnp.asarray(points))[..., 0] error = float(jnp.sum(weights * (jnp.asarray(got) - want) ** 2)) return float(np.sqrt(error / float(jnp.sum(weights * want**2)))) # ── Pictures ──────────────────────────────────────────────────────────────── def draw_field_on_surface(axes, basis, variables, title, keep=None, samples=6): """The projection, painted on the surface of the cut solid. Each face is drawn as its own patch: it has a reference square, so a grid on it maps to a curved quadrilateral in space AND gives the unit-cell preimages the basis wants. The cut is what makes this worth looking at -- the exposed poloidal section shows the blob inside, not just on the skin. """ mesh = basis.mesh grid = np.linspace(0.0, 1.0, samples) tangential = jnp.asarray( np.stack(np.meshgrid(grid, grid, indexing="ij"), -1).reshape(-1, 2) ) faces = geometry.surface_of(mesh, keep) reference = jax.vmap(mesh._face_reference_points, in_axes=(0, None))( jnp.asarray(mesh.faces_left_edge)[faces], tangential ) physical, values = evaluate( basis, variables, jnp.asarray(mesh.faces_left)[faces], reference ) low, high = float(values.min()), float(values.max()) colours = plt.get_cmap("turbo") for patch, patch_values in zip(physical, values): shaped = patch.reshape(samples, samples, 3) axes.plot_surface( shaped[..., 0], shaped[..., 1], shaped[..., 2], facecolors=colours( (patch_values.reshape(samples, samples) - low) / (high - low) ), shade=False, rstride=1, cstride=1, linewidth=0, antialiased=False, ) geometry.tidy_3d(axes, title) def draw_poloidal_slice( axes, basis, variables, layer, toroidal_cells, title, samples=8, levels=(0.0, 1.0) ): """The projection on one poloidal plane, in the section's own coordinates. No cell search anywhere: on a product mesh the plane ``phi = phi_j`` is a face of every cell of toroidal layer ``j``, so it is reached by fixing the third REFERENCE coordinate. The cut is exact, and it costs nothing. """ mesh = basis.mesh grid = np.linspace(0.0, 1.0, samples) square = np.stack(np.meshgrid(grid, grid, indexing="ij"), -1) plane = np.concatenate([square, np.zeros(square.shape[:-1] + (1,))], -1) plane = jnp.asarray(plane.reshape(-1, 3)) # A is the slowest factor, so cell (a, b) is `a * n_b + b`: layer `b` is # every `toroidal_cells`-th cell. n_section = mesh.n_cells_total // toroidal_cells cells = jnp.arange(n_section) * toroidal_cells + layer physical, values = evaluate( basis, variables, cells, jnp.broadcast_to(plane, (n_section,) + plane.shape) ) physical = physical.reshape(n_section, samples, samples, 3) values = values.reshape(n_section, samples, samples) # Back to the machine's own (R, Z) plane, whatever the section was. radii = np.hypot(physical[..., 0], physical[..., 1]) heights = physical[..., 2] # ⚠ One scale for every slice, fixed by the caller. Normalising each panel # on its own maximum makes a blob that has decayed by an order of magnitude # look as bright as the one at its centre -- the picture would then be about # the colourmap rather than about the field. low, high = levels for radius, height, patch in zip(radii, heights, values): axes.pcolormesh( radius, height, patch, cmap="turbo", vmin=low, vmax=high, shading="gouraud" ) axes.plot(radius[0], height[0], color="0.3", linewidth=0.4) axes.plot(radius[-1], height[-1], color="0.3", linewidth=0.4) axes.plot(radius[:, 0], height[:, 0], color="0.3", linewidth=0.4) axes.plot(radius[:, -1], height[:, -1], color="0.3", linewidth=0.4) axes.set_title(title) axes.set_aspect("equal") axes.set_xticks([]) axes.set_yticks([]) return float(values.max()) def run(name, section, toroidal_cells, major_radius, centre, store_mass=False): """Build, project, measure, draw -- for one machine. The circular section and the JET wall go through this unchanged: the only difference between them is which section is revolved and where the perturbation sits. """ print(f"\n=== {name}") started = time.perf_counter() mesh = geometry.revolve(section, toroidal_cells, major_radius) geometry.report(name, mesh, geometry.pappus_volume(section, major_radius)) print(f" maillage {time.perf_counter() - started:6.2f} s") field = make_blob(centre) basis = make_basis(mesh) started = time.perf_counter() variables = VariablesDG(basis=basis, nb_variables=1, store_mass=store_mass) if store_mass: jax.block_until_ready(variables._mass_factorisation[1]) print(f" masses stockees {time.perf_counter() - started:6.2f} s") # ⚠ `project_jit`, not `project`. The projector is a `vmap` over the cells # either way, but un-jitted it runs op by op: measured on a 240-cell # tokamak, 1.68 s eager against 0.22 s for the first jitted call -- and the # compile is paid once for any number of projections, which is what a time # loop does. started = time.perf_counter() variables.project_jit(field) # Finished HERE, not later inside the plotting -- see the module docstring. jax.block_until_ready(variables.dofsl) first = time.perf_counter() - started started = time.perf_counter() variables.project_jit(field) jax.block_until_ready(variables.dofsl) n_dofs = mesh.n_cells_total * (BASIS_ORDER + 1) ** mesh.dim print( f" projection Q{BASIS_ORDER} {first:6.2f} s compile+calcul, " f"{time.perf_counter() - started:6.3f} s ensuite ({n_dofs} ddl)" ) started = time.perf_counter() error = projection_error(mesh, basis, variables, field) print( f" erreur L2 relative {error:9.3e} ({time.perf_counter() - started:.2f} s)" ) figure = plt.figure(figsize=(15, 4.6)) draw_field_on_surface( figure.add_subplot(141, projection="3d"), basis, variables, "sur la coupe", keep=geometry.quarter_cut, ) peaks = [] for panel, layer in enumerate((0, toroidal_cells // 8, toroidal_cells // 4)): peak = draw_poloidal_slice( figure.add_subplot(1, 4, 2 + panel), basis, variables, layer, toroidal_cells, "", ) peaks.append(peak) # The maximum ON THAT PLANE, which the shared scale is what lets one # read off the picture rather than guess. figure.axes[-1].set_title( f"phi = {2 * layer / toroidal_cells:.2f} pi (max {peak:.3f})" ) # A projection cannot exceed the maximum of what it projects, so anything # above 1 measures under-resolution. Reported rather than left in the # picture, where it would just look like a brighter blob. print( f" max sur phi = 0 {peaks[0]:9.3f}" f" (la gaussienne plafonne a 1 : depassement {peaks[0] - 1.0:+.1%})" ) figure.suptitle(f"projection DG d'une perturbation -- {name}") plt.tight_layout() def main(): run( "section circulaire", geometry.poloidal_section(), geometry.N_TOROIDAL, geometry.MAJOR_RADIUS, centre=(geometry.MAJOR_RADIUS + BLOB_OFFSET, 0.0), ) # The JET wall is given in machine (R, Z), so the map adds nothing and the # perturbation is placed at real coordinates: outboard of the magnetic axis, # on the midplane. run( "paroi JET", geometry.jet_section(), geometry.JET_TOROIDAL, 0.0, centre=(3.2, 0.3), store_mass=True, ) plt.show() if __name__ == "__main__": main()