diff --git a/essos/dynamics.py b/essos/dynamics.py index 083d334..74587d9 100644 --- a/essos/dynamics.py +++ b/essos/dynamics.py @@ -1497,6 +1497,23 @@ def loss_fraction(self,r_max=1.0): total_particles_lost = loss_fractions[-1] * len(self.trajectories) return loss_fractions, total_particles_lost, lost_times + def soft_loss_fraction(self, r_max=0.99, width=0.02): + """Differentiable surrogate for the final value of `loss_fraction`. + + Each particle contributes ``sigmoid((r_soft - r_max) / width)``, where ``r_soft`` + is the softmax-weighted mean of its radial coordinate over time, a soft maximum + that spreads the gradient over every sample near the radial excursion peak; + `jnp.max` would route the gradient through the single peak sample and leave the + objective almost flat. It is preferred over a plain ``width * logsumexp(r/width)``, + which carries a ``width * log(n_times)`` offset set by the save grid rather than + by the orbit. The surrogate converges to `loss_fraction` as `width` goes to zero. + Saves past a terminating event count as reaching the boundary. + """ + trajectories_r = self.trajectories[:, :, 0] + trajectories_r = jnp.where(jnp.isfinite(trajectories_r), trajectories_r, 1.0) + r_soft = jnp.sum(trajectories_r * jax.nn.softmax(trajectories_r / width, axis=1), axis=1) + return jnp.mean(jax.nn.sigmoid((r_soft - r_max) / width)) + @partial(jit, static_argnums=(0,1)) diff --git a/essos/fields.py b/essos/fields.py index f1b2a27..2d3fbab 100644 --- a/essos/fields.py +++ b/essos/fields.py @@ -306,8 +306,11 @@ def _radial_interp(s, grid, table, xm, covariant_s=False, half_grid=False, axis_ powers = jnp.stack([1 / q, jnp.ones_like(q), q, q * q, q * q * q]) # q**(2 p) for 2 p = -1..3 return (powers @ (k == np.arange(-1, 4)[:, None])) * ((1 - t) * scaled[i] + t * scaled[i + 1]) +VMEC_WOUT_ARRAYS = ('bmnc', 'xm', 'xn', 'rmnc', 'zmns', 'bsubsmns', 'bsubumnc', 'bsubvmnc', + 'bsupumnc', 'bsupvmnc', 'gmnc', 'xm_nyq', 'xn_nyq', 'Aminor_p') + class Vmec(): - """VMEC equilibrium from a wout file. + """VMEC equilibrium from a wout file, or from its arrays with ``Vmec.from_arrays``. ``mode_tolerance`` drops a Fourier mode when, in every table of its set, its largest amplitude over the radial grid is below that fraction of the @@ -319,24 +322,54 @@ def __init__(self, wout_filename, ntheta=50, nphi=50, close=True, range_torus='f self.wout_filename = wout_filename from netCDF4 import Dataset self.nc = Dataset(self.wout_filename) - self.nfp = int(self.nc.variables["nfp"][0]) - self.bmnc = jnp.array(self.nc.variables["bmnc"][:]) - self.xm = jnp.array(self.nc.variables["xm"][:]) - self.xn = jnp.array(self.nc.variables["xn"][:]) - self.rmnc = jnp.array(self.nc.variables["rmnc"][:]) - self.zmns = jnp.array(self.nc.variables["zmns"][:]) - self.bsubsmns = jnp.array(self.nc.variables["bsubsmns"][:]) - self.bsubumnc = jnp.array(self.nc.variables["bsubumnc"][:]) - self.bsubvmnc = jnp.array(self.nc.variables["bsubvmnc"][:]) - self.bsupumnc = jnp.array(self.nc.variables["bsupumnc"][:]) - self.bsupvmnc = jnp.array(self.nc.variables["bsupvmnc"][:]) - self.gmnc = jnp.array(self.nc.variables["gmnc"][:]) - self.xm_nyq = jnp.array(self.nc.variables["xm_nyq"][:]) - self.xn_nyq = jnp.array(self.nc.variables["xn_nyq"][:]) + self._set_state(nfp=int(self.nc.variables["nfp"][0]), ns=int(self.nc.variables["ns"][0]), + ntheta=ntheta, nphi=nphi, close=close, range_torus=range_torus, + mode_tolerance=mode_tolerance, + **{name: jnp.array(self.nc.variables[name][:]) for name in VMEC_WOUT_ARRAYS}) + + @classmethod + def from_arrays(cls, nfp, ns, bmnc, xm, xn, rmnc, zmns, bsubsmns, bsubumnc, bsubvmnc, + bsupumnc, bsupvmnc, gmnc, xm_nyq, xn_nyq, Aminor_p, + ntheta=50, nphi=50, close=True, range_torus='full torus', mode_tolerance=0.0): + """Build a Vmec field from wout quantities held in memory. + + The arguments carry the wout variable names and are stored as given, so JAX + tracers reach B, AbsB and the traced trajectories and the field stays + differentiable with respect to its spectral coefficients. ``nfp``, ``ns`` and + the mode numbers set array shapes and must be concrete, and so must the + tables when ``mode_tolerance`` is positive, since it selects modes by amplitude. + """ + self = cls.__new__(cls) + self.wout_filename = None + self.nc = None + self._set_state(nfp=nfp, ns=ns, bmnc=bmnc, xm=xm, xn=xn, rmnc=rmnc, zmns=zmns, + bsubsmns=bsubsmns, bsubumnc=bsubumnc, bsubvmnc=bsubvmnc, + bsupumnc=bsupumnc, bsupvmnc=bsupvmnc, gmnc=gmnc, xm_nyq=xm_nyq, + xn_nyq=xn_nyq, Aminor_p=Aminor_p, ntheta=ntheta, nphi=nphi, + close=close, range_torus=range_torus, mode_tolerance=mode_tolerance) + return self + + def _set_state(self, nfp, ns, bmnc, xm, xn, rmnc, zmns, bsubsmns, bsubumnc, bsubvmnc, + bsupumnc, bsupvmnc, gmnc, xm_nyq, xn_nyq, Aminor_p, + ntheta, nphi, close, range_torus, mode_tolerance=0.0): + self.nfp = nfp + self.bmnc = bmnc + self.xm = xm + self.xn = xn + self.rmnc = rmnc + self.zmns = zmns + self.bsubsmns = bsubsmns + self.bsubumnc = bsubumnc + self.bsubvmnc = bsubvmnc + self.bsupumnc = bsupumnc + self.bsupvmnc = bsupvmnc + self.gmnc = gmnc + self.xm_nyq = xm_nyq + self.xn_nyq = xn_nyq if mode_tolerance > 0: self._drop_small_modes(mode_tolerance) self.len_xm_nyq = len(self.xm_nyq) - self.ns = self.nc.variables["ns"][0] + self.ns = ns self.s_full_grid = jnp.linspace(0, 1, self.ns) self.ds = self.s_full_grid[1] - self.s_full_grid[0] self.s_half_grid = self.s_full_grid[1:] - 0.5 * self.ds @@ -346,9 +379,9 @@ def __init__(self, wout_filename, ntheta=50, nphi=50, close=True, range_torus='f self.ntor = int(jnp.max(jnp.abs(self.xn)) / self.nfp) self.range_torus = range_torus self._surface = SurfaceRZFourier.from_vmec(self, ntheta=ntheta, nphi=nphi, close=close, range_torus=range_torus) - self.Aminor_p = jnp.array(self.nc.variables["Aminor_p"][:]) + self.Aminor_p = Aminor_p #self._classifier=SurfaceClassifier(self._surface,p=1,h=0.05) - + def _drop_small_modes(self, tolerance): for tables, numbers in ((('rmnc', 'zmns'), ('xm', 'xn')), (('bmnc', 'gmnc', 'bsubsmns', 'bsubumnc', 'bsubvmnc', 'bsupumnc', 'bsupvmnc'), diff --git a/essos/objective_functions.py b/essos/objective_functions.py index 8aa40a4..ba03317 100644 --- a/essos/objective_functions.py +++ b/essos/objective_functions.py @@ -163,6 +163,16 @@ def loss_particle_iota(field, particles, timestep=1.e-8, maxtime=1e-5, num_steps B_phi=jnp.multiply(B_particle[:,:,0],dphi_dx)+jnp.multiply(B_particle[:,:,1],dphi_dy) return jnp.sum(jnp.maximum(target_iota-B_theta/B_phi,0.0)) +def loss_soft_lost_fraction(field, particles, timestep=1.e-8, maxtime=1e-5, num_steps=300, trace_tolerance=1e-5, model='GuidingCenterAdaptative',boundary=None,r_max=0.99,width=0.02): + """Smooth surrogate of the fraction of particles lost past ``r_max``. + + Differentiable with respect to the field, so alpha confinement can drive + gradient-based optimization. Report `Tracing.loss_fraction` for the exact value. + """ + tracing = Tracing(field=field, model=model, particles=particles, maxtime=maxtime, + timestep=timestep,times_to_trace=num_steps, atol=trace_tolerance,rtol=trace_tolerance,boundary=boundary) + return tracing.soft_loss_fraction(r_max=r_max, width=width) + diff --git a/tests/test_dynamics.py b/tests/test_dynamics.py index b328665..71aa370 100644 --- a/tests/test_dynamics.py +++ b/tests/test_dynamics.py @@ -1,3 +1,4 @@ +import os import pytest import numpy as np from pathlib import Path @@ -23,9 +24,12 @@ _VMEC_GUIDING_CENTER_MODELS, ) from essos.background_species import BackgroundSpecies -from essos.fields import Vmec +from essos.fields import Vmec, VMEC_WOUT_ARRAYS from essos.surfaces import SurfaceClassifier, SurfaceRZFourier +WOUT_FILE = os.path.join(os.path.dirname(__file__), "..", "examples", "input_files", + "wout_LandremanPaul2021_QA_reactorScale_lowres.nc") + def test_particles_initialization_all_params(): nparticles = 100 initial_xyz = jnp.array([[1.0, 0.0, 0.0]] * nparticles) @@ -449,6 +453,58 @@ def test_vmec_collisional_guiding_centers_cross_the_magnetic_axis(model): assert (s[:, -1] > 1.5 * s.min(axis=1)).all() +def radial_tracing(peaks, times=jnp.linspace(0.0, 1.0, 40)): + """Tracing holding prescribed radial excursions, bypassing the ODE solve.""" + trajectories_r = 0.5 + (peaks[:, None] - 0.5) * jnp.sin(jnp.pi * times)[None, :] + tracing = Tracing.__new__(Tracing) + tracing.trajectories = jnp.stack([trajectories_r, jnp.zeros_like(trajectories_r), jnp.zeros_like(trajectories_r)], axis=-1) + tracing.times = times + return tracing + +def test_soft_loss_fraction_converges_to_loss_fraction(): + tracing = radial_tracing(jnp.array([0.70, 0.93, 0.86, 0.99])) + exact = tracing.loss_fraction(r_max=0.9)[0][-1] + assert exact == 0.5 + + errors = [abs(float(tracing.soft_loss_fraction(r_max=0.9, width=width) - exact)) for width in (0.02, 0.01, 0.005, 0.002)] + assert errors == sorted(errors, reverse=True) + assert errors[-1] < 1e-3 + +def test_soft_loss_fraction_gradient_is_nonzero_where_loss_fraction_is_flat(): + peaks = jnp.array([0.70, 0.93, 0.86, 0.99]) + exact_gradient = jax.grad(lambda p: radial_tracing(p).loss_fraction(r_max=0.9)[0][-1])(peaks) + soft_gradient = jax.grad(lambda p: radial_tracing(p).soft_loss_fraction(r_max=0.9, width=0.01))(peaks) + + assert jnp.all(exact_gradient == 0.0) + assert jnp.all(soft_gradient[1:] > 0.0) + +def vmec_alpha_tracing(field, nparticles=4, maxtime=4e-6, times_to_trace=10): + theta = jnp.linspace(0, 2*jnp.pi, nparticles) + phi = jnp.linspace(0, 2*jnp.pi/field.nfp, nparticles) + particles = Particles(initial_xyz=jnp.array([0.85*jnp.ones(nparticles), theta, phi]).T, mass=ALPHA_PARTICLE_MASS, + charge=ALPHA_PARTICLE_CHARGE, energy=FUSION_ALPHA_PARTICLE_ENERGY, field=field) + return Tracing(field=field, model='GuidingCenterAdaptative', particles=particles, maxtime=maxtime, + timestep=1e-8, times_to_trace=times_to_trace, atol=1e-5, rtol=1e-5) + +def test_vmec_from_arrays_traces_identically(): + vmec = Vmec(WOUT_FILE) + rebuilt = Vmec.from_arrays(nfp=vmec.nfp, ns=vmec.ns, **{name: getattr(vmec, name) for name in VMEC_WOUT_ARRAYS}) + + assert jnp.array_equal(vmec_alpha_tracing(rebuilt).trajectories, vmec_alpha_tracing(vmec).trajectories) + +def test_soft_loss_fraction_differentiates_vmec_coefficients(): + vmec = Vmec(WOUT_FILE) + arrays = {name: getattr(vmec, name) for name in VMEC_WOUT_ARRAYS} + scaled = ('bmnc', 'bsubsmns', 'bsubumnc', 'bsubvmnc', 'bsupumnc', 'bsupvmnc') + + def soft_loss_of_field_scale(scale): + field = Vmec.from_arrays(nfp=vmec.nfp, ns=vmec.ns, + **{**arrays, **{name: arrays[name]*scale for name in scaled}}) + return vmec_alpha_tracing(field).soft_loss_fraction(r_max=0.88, width=0.01) + + assert jax.grad(soft_loss_of_field_scale)(1.0) != 0.0 + + def test_tracing_initialization(field, particles,electric_field): x = jnp.linspace(1, 2, particles.nparticles) y = jnp.zeros(particles.nparticles) diff --git a/tests/test_fields.py b/tests/test_fields.py index 5f0f380..a5b0ef7 100644 --- a/tests/test_fields.py +++ b/tests/test_fields.py @@ -1,10 +1,14 @@ +import os import pytest from pathlib import Path from essos.coils import Coils, Curves -from essos.fields import BiotSavart +from essos.fields import BiotSavart, Vmec, VMEC_WOUT_ARRAYS import jax import jax.numpy as jnp -from jax import random +from jax import random, vmap + +WOUT_FILE = os.path.join(os.path.dirname(__file__), "..", "examples", "input_files", + "wout_LandremanPaul2021_QA_reactorScale_lowres.nc") class MockCoils: def __init__(self): @@ -82,6 +86,27 @@ def test_biot_savart_cylindrical_interface_matches_cartesian_and_differentiates( # dAbsB_by_dX = biot_savart.dAbsB_by_dX(points) # assert jnp.allclose(dAbsB_by_dX, jnp.array([7.16688661e-05, 3.82872752e-05, 1.01490560e-04])) +def test_vmec_from_arrays_matches_wout_file(): + vmec = Vmec(WOUT_FILE) + rebuilt = Vmec.from_arrays(nfp=vmec.nfp, ns=vmec.ns, + **{name: getattr(vmec, name) for name in VMEC_WOUT_ARRAYS}) + points = jnp.array([[0.3, 0.4, 0.5], [0.7, 1.2, 0.2], [0.9, 3.0, 1.1]]) + + assert (rebuilt.nfp, rebuilt.ns, rebuilt.mpol, rebuilt.ntor) == (vmec.nfp, vmec.ns, vmec.mpol, vmec.ntor) + assert jnp.array_equal(vmap(rebuilt.B)(points), vmap(vmec.B)(points)) + assert jnp.array_equal(vmap(rebuilt.AbsB)(points), vmap(vmec.AbsB)(points)) + assert jnp.array_equal(rebuilt.surface.gamma, vmec.surface.gamma) + +def test_vmec_from_arrays_is_differentiable_in_the_coefficients(): + vmec = Vmec(WOUT_FILE) + arrays = {name: getattr(vmec, name) for name in VMEC_WOUT_ARRAYS} + point = jnp.array([0.7, 1.2, 0.2]) + + def AbsB_of_scale(scale): + return Vmec.from_arrays(nfp=vmec.nfp, ns=vmec.ns, **{**arrays, 'bmnc': arrays['bmnc']*scale}).AbsB(point) + + assert jnp.isclose(jax.grad(AbsB_of_scale)(1.0), vmec.AbsB(point)) + if __name__ == "__main__": pytest.main() @@ -171,6 +196,10 @@ def test_vmec_mode_tolerance_keeps_the_field(): for name in ("AbsB", "B_contravariant", "to_xyz"): a, b = jax.vmap(getattr(full, name))(points), jax.vmap(getattr(truncated, name))(points) assert jnp.abs(a - b).max() < 5e-3 * jnp.abs(a).max() + rebuilt = Vmec.from_arrays(nfp=full.nfp, ns=full.ns, ntheta=8, nphi=8, mode_tolerance=1e-3, + **{name: getattr(full, name) for name in VMEC_WOUT_ARRAYS}) + assert jnp.array_equal(rebuilt.xm_nyq, truncated.xm_nyq) and jnp.array_equal(rebuilt.xm, truncated.xm) + assert jnp.array_equal(jax.vmap(rebuilt.AbsB)(points), jax.vmap(truncated.AbsB)(points)) def test_fused_guiding_center_quantities_match_the_separate_methods(): diff --git a/tests/test_objective_functions.py b/tests/test_objective_functions.py index 701ba71..536439d 100644 --- a/tests/test_objective_functions.py +++ b/tests/test_objective_functions.py @@ -9,6 +9,7 @@ import jax.numpy as jnp import essos.objective_functions as objf +from essos.dynamics import Tracing class DummyCoils: @@ -89,6 +90,9 @@ def __init__(self, *args, **kwargs): self.times_to_trace = 4 self.maxtime = 1e-5 + def soft_loss_fraction(self, r_max=0.99, width=0.02): + return Tracing.soft_loss_fraction(self, r_max=r_max, width=width) + @jax.tree_util.register_pytree_node_class class DummySurface: @@ -219,6 +223,8 @@ def test_particle_losses(self, tracing): self.assertTrue(jnp.isfinite(objf.loss_particle_rcross_final(self.field, self.particles))) self.assertTrue(jnp.isfinite(objf.loss_particle_Br(self.field, self.particles))) self.assertTrue(jnp.isfinite(objf.loss_particle_iota(self.field, self.particles))) + soft_lost_fraction = objf.loss_soft_lost_fraction(self.field, self.particles, r_max=1.2, width=0.05) + self.assertTrue(0.0 <= soft_lost_fraction <= 1.0) @patch("essos.objective_functions.BdotN_over_B", return_value=jnp.ones((2, 3), dtype=jnp.float64)) def test_surface_losses(self, bdotn):