From 1a302f68e5f434eed760d297283bb36d65bbe7d2 Mon Sep 17 00:00:00 2001 From: Rogerio Jorge Date: Wed, 19 Aug 2026 14:27:52 -0500 Subject: [PATCH 1/5] Add in-memory Vmec construction and a differentiable loss fraction surrogate Vmec.from_arrays builds the field from wout coefficients already in memory, so a caller holding live JAX arrays no longer has to write a netCDF file to trace in a VMEC equilibrium. Tracing.soft_loss_fraction is a smooth surrogate for the final loss fraction, giving alpha confinement a usable gradient, and loss_soft_lost_fraction exposes it as an objective. The exact loss_fraction diagnostic is untouched. --- essos/dynamics.py | 17 ++++++++ essos/fields.py | 69 +++++++++++++++++++++++-------- essos/objective_functions.py | 10 +++++ tests/test_dynamics.py | 58 +++++++++++++++++++++++++- tests/test_fields.py | 33 ++++++++++++++- tests/test_objective_functions.py | 6 +++ 6 files changed, 172 insertions(+), 21 deletions(-) diff --git a/essos/dynamics.py b/essos/dynamics.py index 083d3341..74587d90 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 f1b2a274..2d3fbab9 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 8aa40a4a..ba033171 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 b3286651..71aa3709 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 5f0f3803..a5b0ef78 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 701ba718..536439d9 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): From 6ea0915e5746a84b47b9180d71062e72e5a8e8ee Mon Sep 17 00:00:00 2001 From: Rogerio Jorge Date: Wed, 30 Sep 2026 15:16:20 -0500 Subject: [PATCH 2/5] Validate differentiable loss inputs and preserve traced initial states --- essos/dynamics.py | 37 ++++++++++++++++--------------- essos/fields.py | 17 ++++++-------- tests/test_dynamics.py | 32 +++++++++++++++++++++++++- tests/test_fields.py | 13 +++++++++-- tests/test_objective_functions.py | 6 ++++- 5 files changed, 73 insertions(+), 32 deletions(-) diff --git a/essos/dynamics.py b/essos/dynamics.py index 74587d90..331d3720 100644 --- a/essos/dynamics.py +++ b/essos/dynamics.py @@ -1296,9 +1296,11 @@ def save_interval(carry, _): if len(self.stopping_criteria) > 1: event_sharding = tuple(sharding_index for _ in self.stopping_criteria) output_sharding = (sharding, event_sharding) + initial_conditions = self.initial_conditions + if not isinstance(initial_conditions, jax.core.Tracer): + initial_conditions = np.asarray(jax.device_get(initial_conditions)) if sharding is not None: - initial_conditions = device_put( - np.asarray(jax.device_get(self.initial_conditions)), sharding) + initial_conditions = device_put(initial_conditions, sharding) random_keys = self.particles.random_keys if self.particles else None if random_keys is not None: random_keys = device_put(jax.device_get(random_keys), sharding_index) @@ -1306,8 +1308,7 @@ def save_interval(carry, _): initial_conditions, random_keys) else: device = devices[0] - initial_conditions = device_put( - np.asarray(jax.device_get(self.initial_conditions)), device) + initial_conditions = device_put(initial_conditions, device) random_keys = self.particles.random_keys if self.particles else None if random_keys is not None: random_keys = device_put(jax.device_get(random_keys), device) @@ -1498,21 +1499,21 @@ def loss_fraction(self,r_max=1.0): 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. + """Soft peak-flux crossing score for VMEC guiding centres. + + Boundary stops count as exits; unrelated failures return NaN. A smaller + positive width sharpens the score. Hard losses remain the diagnostic. """ - 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)) + if not isinstance(self.field, Vmec) or self.model not in _VMEC_GUIDING_CENTER_MODELS: + raise ValueError("soft_loss_fraction requires VMEC guiding centres") + if not np.isfinite(width) or width <= 0: + raise ValueError("width must be finite and positive") + radial = self.trajectories[:, :, 0] + finite = jnp.all(jnp.isfinite(self.trajectories), axis=-1) + stopped = (self.boundary_hits & self._has_boundary_event)[:, None] + radial = jnp.where(finite, radial, jnp.where(stopped, 1.0, jnp.nan)) + peak = jnp.sum(radial * jax.nn.softmax(radial / width, axis=1), axis=1) + return jnp.mean(jax.nn.sigmoid((peak - r_max) / width)) diff --git a/essos/fields.py b/essos/fields.py index 2d3fbab9..acf8d41a 100644 --- a/essos/fields.py +++ b/essos/fields.py @@ -331,18 +331,15 @@ def __init__(self, wout_filename, ntheta=50, nphi=50, close=True, range_torus='f 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. + """Build a differentiable VMEC field from in-memory wout arrays. - 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. + Metadata and mode numbers must be concrete. Positive ``mode_tolerance`` + also requires concrete coefficient tables to select 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, + self._set_state(nfp=int(nfp), ns=int(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, @@ -375,8 +372,9 @@ def _set_state(self, nfp, ns, bmnc, xm, xn, rmnc, zmns, bsubsmns, bsubumnc, bsub self.s_half_grid = self.s_full_grid[1:] - 0.5 * self.ds self.r_axis = self.rmnc[0, 0] self.z_axis=self.zmns[0,0] - self.mpol = int(jnp.max(self.xm)) - self.ntor = int(jnp.max(jnp.abs(self.xn)) / self.nfp) + with jax.ensure_compile_time_eval(): + self.mpol = int(jnp.max(self.xm)) + 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 = Aminor_p @@ -609,4 +607,3 @@ def _tree_unflatten(cls, aux_data, children): tree_util.register_pytree_node(CombinedField, CombinedField._tree_flatten, CombinedField._tree_unflatten) - diff --git a/tests/test_dynamics.py b/tests/test_dynamics.py index 71aa3709..13fb41cd 100644 --- a/tests/test_dynamics.py +++ b/tests/test_dynamics.py @@ -459,6 +459,10 @@ def radial_tracing(peaks, times=jnp.linspace(0.0, 1.0, 40)): 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 + tracing.field = Vmec.__new__(Vmec) + tracing.model = "GuidingCenterAdaptative" + tracing._has_boundary_event = True + tracing.boundary_hits = jnp.zeros(len(peaks), dtype=bool) return tracing def test_soft_loss_fraction_converges_to_loss_fraction(): @@ -478,6 +482,29 @@ def test_soft_loss_fraction_gradient_is_nonzero_where_loss_fraction_is_flat(): assert jnp.all(exact_gradient == 0.0) assert jnp.all(soft_gradient[1:] > 0.0) +@pytest.mark.parametrize("width", [0.0, -0.02, float("nan"), float("inf")]) +def test_soft_loss_rejects_invalid_width(width): + with pytest.raises(ValueError, match="width"): + radial_tracing(jnp.array([0.8])).soft_loss_fraction(width=width) + + +def test_soft_loss_distinguishes_boundary_stop_from_failure(): + trace = radial_tracing(jnp.array([0.8])) + trace.trajectories = trace.trajectories.at[0, -1, 1].set(jnp.nan) + assert jnp.isnan(trace.soft_loss_fraction()) + trace.boundary_hits = jnp.array([True]) + assert jnp.isfinite(trace.soft_loss_fraction()) + assert trace.soft_loss_fraction(width=0.001) > 0.99 + + +@pytest.mark.parametrize("model, field", [("Lorentz", Vmec.__new__(Vmec)), ("GuidingCenterAdaptative", object())]) +def test_soft_loss_rejects_nonflux_trajectories(model, field): + trace = radial_tracing(jnp.array([0.8])) + trace.model, trace.field = model, field + with pytest.raises(ValueError, match="VMEC guiding"): + trace.soft_loss_fraction() + + 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) @@ -502,7 +529,10 @@ def soft_loss_of_field_scale(scale): **{**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 + evaluate = jax.jit(jax.value_and_grad(soft_loss_of_field_scale)) + for scale in (1.0, 1.01): + value, gradient = evaluate(scale) + assert jnp.isfinite(value) and jnp.isfinite(gradient) and gradient != 0.0 def test_tracing_initialization(field, particles,electric_field): diff --git a/tests/test_fields.py b/tests/test_fields.py index a5b0ef78..b1c3ea71 100644 --- a/tests/test_fields.py +++ b/tests/test_fields.py @@ -1,4 +1,5 @@ import os +import numpy as np import pytest from pathlib import Path from essos.coils import Coils, Curves @@ -88,7 +89,7 @@ def test_biot_savart_cylindrical_interface_matches_cartesian_and_differentiates( def test_vmec_from_arrays_matches_wout_file(): vmec = Vmec(WOUT_FILE) - rebuilt = Vmec.from_arrays(nfp=vmec.nfp, ns=vmec.ns, + rebuilt = Vmec.from_arrays(nfp=np.int64(vmec.nfp), ns=jnp.asarray(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]]) @@ -102,10 +103,18 @@ def test_vmec_from_arrays_is_differentiable_in_the_coefficients(): arrays = {name: getattr(vmec, name) for name in VMEC_WOUT_ARRAYS} point = jnp.array([0.7, 1.2, 0.2]) + traced = [] + def AbsB_of_scale(scale): + traced.append(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)) + evaluate = jax.jit(jax.value_and_grad(AbsB_of_scale)) + for scale in (1.0, 1.1): + value, gradient = evaluate(scale) + assert jnp.isclose(value, scale * vmec.AbsB(point)) + assert jnp.isclose(gradient, vmec.AbsB(point)) + assert len(traced) == 1 if __name__ == "__main__": pytest.main() diff --git a/tests/test_objective_functions.py b/tests/test_objective_functions.py index 536439d9..331b111e 100644 --- a/tests/test_objective_functions.py +++ b/tests/test_objective_functions.py @@ -10,6 +10,7 @@ import essos.objective_functions as objf from essos.dynamics import Tracing +from essos.fields import Vmec class DummyCoils: @@ -89,6 +90,9 @@ def __init__(self, *args, **kwargs): self.loss_fractions = jnp.array([0.1, 0.2, 1.0]) self.times_to_trace = 4 self.maxtime = 1e-5 + self.model = kwargs.get("model", "GuidingCenterAdaptative") + self.boundary_hits = jnp.zeros(2, dtype=bool) + self._has_boundary_event = True def soft_loss_fraction(self, r_max=0.99, width=0.02): return Tracing.soft_loss_fraction(self, r_max=r_max, width=width) @@ -223,7 +227,7 @@ 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) + soft_lost_fraction = objf.loss_soft_lost_fraction(Vmec.__new__(Vmec), 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)) From 8e6686bf6712e283eed9957d1c1a5ceaf9e9925e Mon Sep 17 00:00:00 2001 From: Rogerio Jorge Date: Wed, 30 Sep 2026 15:25:58 -0500 Subject: [PATCH 3/5] Describe the loss objective as a flux-peak score --- essos/objective_functions.py | 5 ++--- 1 file changed, 2 insertions(+), 3 deletions(-) diff --git a/essos/objective_functions.py b/essos/objective_functions.py index ba033171..16fbe087 100644 --- a/essos/objective_functions.py +++ b/essos/objective_functions.py @@ -164,10 +164,9 @@ def loss_particle_iota(field, particles, timestep=1.e-8, maxtime=1e-5, num_steps 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 flux-peak crossing score for VMEC guiding centres. - Differentiable with respect to the field, so alpha confinement can drive - gradient-based optimization. Report `Tracing.loss_fraction` for the exact value. + Report ``Tracing.loss_fraction`` for hard losses. """ 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) From db6d2c574292de3a7cf5af5b533bb1fbdaf6b969 Mon Sep 17 00:00:00 2001 From: Rogerio Jorge Date: Thu, 1 Oct 2026 16:18:21 -0500 Subject: [PATCH 4/5] Count finite boundary stops in the soft loss score --- essos/dynamics.py | 1 + tests/test_dynamics.py | 10 +++++++--- 2 files changed, 8 insertions(+), 3 deletions(-) diff --git a/essos/dynamics.py b/essos/dynamics.py index 331d3720..20f95f51 100644 --- a/essos/dynamics.py +++ b/essos/dynamics.py @@ -1513,6 +1513,7 @@ def soft_loss_fraction(self, r_max=0.99, width=0.02): stopped = (self.boundary_hits & self._has_boundary_event)[:, None] radial = jnp.where(finite, radial, jnp.where(stopped, 1.0, jnp.nan)) peak = jnp.sum(radial * jax.nn.softmax(radial / width, axis=1), axis=1) + peak = jnp.where(stopped[:, 0], jnp.maximum(peak, 1.0), peak) return jnp.mean(jax.nn.sigmoid((peak - r_max) / width)) diff --git a/tests/test_dynamics.py b/tests/test_dynamics.py index 13fb41cd..0432c51b 100644 --- a/tests/test_dynamics.py +++ b/tests/test_dynamics.py @@ -488,10 +488,14 @@ def test_soft_loss_rejects_invalid_width(width): radial_tracing(jnp.array([0.8])).soft_loss_fraction(width=width) -def test_soft_loss_distinguishes_boundary_stop_from_failure(): +@pytest.mark.parametrize("nonfinite", [False, True]) +def test_soft_loss_distinguishes_boundary_stop_from_failure(nonfinite): trace = radial_tracing(jnp.array([0.8])) - trace.trajectories = trace.trajectories.at[0, -1, 1].set(jnp.nan) - assert jnp.isnan(trace.soft_loss_fraction()) + if nonfinite: + trace.trajectories = trace.trajectories.at[0, -1, 1].set(jnp.nan) + assert jnp.isnan(trace.soft_loss_fraction()) + else: + assert trace.soft_loss_fraction(width=0.001) < 0.01 trace.boundary_hits = jnp.array([True]) assert jnp.isfinite(trace.soft_loss_fraction()) assert trace.soft_loss_fraction(width=0.001) > 0.99 From 1550979ee71c650870bae37f74bf7a1a9ebd927e Mon Sep 17 00:00:00 2001 From: Rogerio Jorge Date: Thu, 1 Oct 2026 17:27:43 -0500 Subject: [PATCH 5/5] Verify sine coefficient gradients through Boozer fields --- tests/test_boozer.py | 7 +++++++ 1 file changed, 7 insertions(+) diff --git a/tests/test_boozer.py b/tests/test_boozer.py index 27db9d80..5ad0c415 100644 --- a/tests/test_boozer.py +++ b/tests/test_boozer.py @@ -302,6 +302,13 @@ def test_sine_orbits_match_rotated_cosine_field(): bmns = np.stack([np.zeros_like(s), -B0*EPS*np.sqrt(s)*np.sin(shift)]) field = BoozerField.from_booz(s, bmnc, [0, 1], [0, 0], np.full_like(s, IOTA), np.full_like(s, G), np.zeros_like(s), PSI0, 1, bmns=bmns) + def value(amplitude): + live = eqx.tree_at(lambda f: f.sine_coef, field, amplitude * field.sine_coef) + return live.modB(0.3, 0.9, 0.0) + expected = -B0 * EPS * np.sqrt(0.3) * np.sin(shift) * np.sin(0.9) + derivative = jax.jit(jax.grad(value)) + for amplitude in (0.0, 1.0): + assert float(derivative(amplitude)) == pytest.approx(expected, rel=1e-12) theta, pitch = np.linspace(0, 6, 8), np.linspace(-0.9, 0.9, 8) kwargs = dict(speed=V0, mass=M, charge=Q, tmax=2e-5, timestep=2e-8, n_save=7) reference = trace_boozer(tokamak(), np.full(8, 0.3), theta, np.zeros(8), pitch, **kwargs)