Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 17 additions & 0 deletions essos/dynamics.py
Original file line number Diff line number Diff line change
Expand Up @@ -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))
Expand Down
69 changes: 51 additions & 18 deletions essos/fields.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand All @@ -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'),
Expand Down
10 changes: 10 additions & 0 deletions essos/objective_functions.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)




Expand Down
58 changes: 57 additions & 1 deletion tests/test_dynamics.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import os
import pytest
import numpy as np
from pathlib import Path
Expand All @@ -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)
Expand Down Expand Up @@ -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)
Expand Down
33 changes: 31 additions & 2 deletions tests/test_fields.py
Original file line number Diff line number Diff line change
@@ -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):
Expand Down Expand Up @@ -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()

Expand Down Expand Up @@ -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():
Expand Down
6 changes: 6 additions & 0 deletions tests/test_objective_functions.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@
import jax.numpy as jnp

import essos.objective_functions as objf
from essos.dynamics import Tracing


class DummyCoils:
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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):
Expand Down
Loading