Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
83 commits
Select commit Hold shift + click to select a range
428ac62
Fix bugs in compute quantities
unalmis May 19, 2026
345384d
Apply suggestions from code review
unalmis May 19, 2026
508663c
Apply suggestions from code review
unalmis May 19, 2026
d5cc10b
Merge branch 'ku/available_energy' into ku/compute_bugs
unalmis May 19, 2026
f6253bf
Merge branch 'ku/available_energy' into ku/compute_bugs
unalmis May 19, 2026
e187315
Fixing flunked merges hopefully
unalmis May 19, 2026
f81d3d2
Merge branch 'ku/available_energy' into ku/compute_bugs
unalmis May 19, 2026
9f9df69
Merge branch 'ku/compute_bugs' into ku/gamma
unalmis May 19, 2026
88b1be9
fixing docs
unalmis May 19, 2026
4d9f7bb
think this fixes everything
unalmis May 19, 2026
8577da4
fix reference links
unalmis May 19, 2026
1d64f6e
clarify docstring
unalmis May 19, 2026
399a280
fix failing doc
unalmis May 19, 2026
26d59ec
make logic clearer for reviewer
unalmis May 20, 2026
9b5a5da
old
unalmis May 20, 2026
6205eca
This commit adds alpha folding.
unalmis May 21, 2026
afcceca
Clean up.
unalmis May 21, 2026
6864aff
update comptue data
unalmis May 21, 2026
d16ee52
clean
unalmis May 21, 2026
a50da70
reduce diff size
unalmis May 21, 2026
9f988d9
Merge branch 'ku/available_energy' into ku/compute_bugs
unalmis May 22, 2026
d7cd5c3
Merge branch 'ku/compute_bugs' into ku/gamma
unalmis May 22, 2026
c59e102
Merge branch 'ku/gamma' into ku/alpha_fold
unalmis May 22, 2026
685b557
add info on flux tube vs orbit models
unalmis May 23, 2026
3ccf13d
fix typo
unalmis May 23, 2026
09061d3
Merge branch 'ku/gamma' into ku/alpha_fold
unalmis May 23, 2026
74cb787
.
unalmis May 23, 2026
db73628
Merge branch 'ku/gamma' into ku/alpha_fold
unalmis May 23, 2026
5089445
update docs
unalmis May 24, 2026
27bae73
Merge branch 'ku/gamma' into ku/alpha_fold
unalmis May 24, 2026
7679a05
Update _turbulence.py
unalmis May 24, 2026
5f6ab12
.
unalmis May 26, 2026
b53c68e
Merge branch 'ku/available_energy' into ku/compute_bugs
unalmis May 26, 2026
f090571
Merge branch 'ku/compute_bugs' into ku/gamma
unalmis May 26, 2026
a648360
Merge branch 'ku/gamma' into ku/alpha_fold
unalmis May 26, 2026
7f1fe7e
.
unalmis May 27, 2026
2433d3c
Merge branch 'ku/gamma' into ku/alpha_fold
unalmis May 27, 2026
7bb6dd8
fix nan
unalmis May 31, 2026
06c7296
.
unalmis May 31, 2026
f8528ba
.
unalmis May 31, 2026
e966db5
Fix references to constants in test_objective_funs
unalmis May 31, 2026
7620912
Merge branch 'ku/gamma' into ku/alpha_fold
unalmis May 31, 2026
054f387
Merge branch 'ku/available_energy' into ku/compute_bugs
unalmis Jun 10, 2026
7af571f
Merge branch 'ku/compute_bugs' into ku/gamma
unalmis Jun 10, 2026
34a71de
Merge branch 'ku/gamma' into ku/alpha_fold
unalmis Jun 10, 2026
58e9618
clean up
unalmis Jun 15, 2026
e838df3
Merge branch 'ku/gamma' into ku/alpha_fold
unalmis Jun 15, 2026
1b79258
Merge branch 'ku/available_energy' into ku/compute_bugs
unalmis Jun 17, 2026
7228ced
Merge branch 'ku/compute_bugs' into ku/gamma
unalmis Jun 17, 2026
3fc3c3e
Merge branch 'ku/gamma' into ku/alpha_fold
unalmis Jun 17, 2026
e0c59e9
Merge branch 'ku/available_energy' into ku/compute_bugs
unalmis Jun 19, 2026
9c93885
Merge remote-tracking branch 'upstream/ku/compute_bugs' into ku/gamma
unalmis Jun 19, 2026
f72466c
merge
unalmis Jun 19, 2026
3a0471f
Merge remote-tracking branch 'upstream/ku/gamma' into ku/alpha_fold
unalmis Jun 19, 2026
d8b3f01
merge
unalmis Jun 19, 2026
36bdef2
Merge branch 'ku/gamma' into ku/alpha_fold
unalmis Jun 19, 2026
bb566ed
Merge branch 'ku/available_energy' into ku/compute_bugs
unalmis Jul 20, 2026
c530f15
Merge branch 'ku/compute_bugs' into ku/gamma
unalmis Jul 20, 2026
53ada21
Merge branch 'ku/gamma' into ku/alpha_fold
unalmis Jul 20, 2026
3fd37df
Merge branch 'ku/available_energy' into ku/compute_bugs
unalmis Jul 20, 2026
c17d522
Merge branch 'ku/compute_bugs' into ku/gamma
unalmis Jul 20, 2026
79bd4e1
Merge branch 'ku/gamma' into ku/alpha_fold
unalmis Jul 20, 2026
04e9f31
Merge branch 'ku/available_energy' into ku/compute_bugs
unalmis Jul 20, 2026
deff24e
Merge branch 'ku/compute_bugs' into ku/gamma
unalmis Jul 20, 2026
0c18bfa
Merge branch 'ku/gamma' into ku/alpha_fold
unalmis Jul 20, 2026
3b9738f
Merge branch 'ku/available_energy' into ku/compute_bugs
unalmis Jul 20, 2026
1ffe586
Merge branch 'ku/compute_bugs' into ku/gamma
unalmis Jul 20, 2026
a353ed0
Merge branch 'ku/gamma' into ku/alpha_fold
unalmis Jul 20, 2026
8f1b1d9
Merge branch 'ku/available_energy' into ku/compute_bugs
unalmis Jul 24, 2026
ae3de60
Merge branch 'ku/compute_bugs' into ku/gamma
unalmis Jul 24, 2026
f01f3bd
Merge branch 'ku/gamma' into ku/alpha_fold
unalmis Jul 24, 2026
954d8fb
Merge branch 'ku/available_energy' into ku/compute_bugs
unalmis Jul 29, 2026
4bac33b
Merge branch 'ku/compute_bugs' into ku/gamma
unalmis Jul 29, 2026
2c088e7
Merge branch 'ku/gamma' into ku/alpha_fold
unalmis Jul 29, 2026
f0f86b3
Merge branch 'ku/available_energy' into ku/compute_bugs
unalmis Jul 31, 2026
866558e
Merge branch 'ku/compute_bugs' into ku/gamma
unalmis Jul 31, 2026
6828091
Merge branch 'ku/gamma' into ku/alpha_fold
unalmis Jul 31, 2026
5014ff4
Merge branch 'ku/available_energy' into ku/alpha_fold
unalmis Aug 1, 2026
9cf71fd
Merge branch 'ku/available_energy' into ku/alpha_fold
unalmis Aug 1, 2026
7a59a9c
Merge branch 'ku/available_energy' into ku/alpha_fold
unalmis Aug 1, 2026
53ab60f
Merge branch 'ku/available_energy' into ku/alpha_fold
unalmis Aug 1, 2026
27bc71b
Merge branch 'ku/available_energy' into ku/alpha_fold
unalmis Aug 1, 2026
8c5693c
Merge branch 'ku/available_energy' into ku/alpha_fold
unalmis Aug 3, 2026
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
150 changes: 122 additions & 28 deletions desc/compute/_fast_ion.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,7 @@

from ..batching import batch_map
from ..integrals.bounce_integral import Bounce2D, Options
from ..integrals.quad_utils import _LossCone
from ..integrals.quad_utils import _LossCone, _periodic_voronoi_widths
from ..utils import cross, dot, safediv
from ._drift import (
_alpha_drift_wb_inverse,
Expand Down Expand Up @@ -249,31 +249,106 @@ def _reduction_gamma_c(v_tau, radial, poloidal, opts=None):
return (v_tau * _gamma_c(radial, poloidal) ** 2).sum(-1).mean(-2)


def _reduction_gamma_delta(v_tau, radial, poloidal, opts):
v_tau = v_tau.mean(-3)
def _reshape_iota(iota, suffix_ndim):
"""Reshape iota to broadcast over prefixed surface data."""
iota = jnp.asarray(iota)
prefix = () if iota.size == 1 else iota.shape
return iota.reshape(prefix + (1,) * suffix_ndim)


def _well_field_period(points, NFP):
"""Locate each long-field-line well by its midpoint field period."""
z1, z2 = points
midpoint = 0.5 * (z1 + z2)
return ((NFP * midpoint) // (2 * jnp.pi)).astype(jnp.int32)


def _local_well_rank(field_period, valid):
"""Rank wells within each single-field-period segment."""
well = jnp.arange(field_period.shape[-1])
previous = well[None, :] < well[:, None]
return (
(field_period[..., None] == field_period[..., None, :])
& valid[..., None, :]
& previous
).sum(-1)


def _fold_wells_to_alpha(values, points, opts, iota, NFP):
"""Treat long-field-line wells as denser alpha samples over one field period."""
z1, z2 = points
valid_well = z1 < z2
well_field_period = _well_field_period(points, NFP)
local_well_index = _local_well_rank(well_field_period, valid_well)
num_alpha, num_pitch, num_well = z1.shape[-3:]
prefix = z1.shape[:-3]

field_period_index = jnp.arange(opts.num_field_periods).reshape(
(1, opts.num_field_periods, 1, 1, 1)
)
local_well_slot = jnp.arange(num_well).reshape((1, 1, 1, 1, num_well))
# Axes: prefix, base alpha, field period, pitch, original well, local well.
well_to_alpha = (
valid_well[..., :, None, :, :, None]
& (well_field_period[..., :, None, :, :, None] == field_period_index)
& (local_well_index[..., :, None, :, :, None] == local_well_slot)
)

def fold(value):
value = jnp.where(well_to_alpha, value[..., :, None, :, :, None], 0.0).sum(-2)
return value.reshape(
prefix + (num_alpha * opts.num_field_periods, num_pitch, num_well)
)

alpha = opts.alpha.reshape((1,) * len(prefix) + (num_alpha, 1))
alpha = alpha + _reshape_iota(iota, 2) * (2 * jnp.pi / NFP) * jnp.arange(
opts.num_field_periods
)
alpha = (alpha % (2 * jnp.pi)).reshape(
prefix + (num_alpha * opts.num_field_periods, 1, 1)
)
mask = well_to_alpha.any(-2).reshape(
prefix + (num_alpha * opts.num_field_periods, num_pitch, num_well)
)
return (*[fold(value) for value in values], alpha, mask)


def _alpha_weights(alpha, valid, period=2 * jnp.pi):
"""Periodic Voronoi cell widths for a possibly nonuniform alpha grid."""
alpha = alpha.swapaxes(-3, -1)
valid = valid.swapaxes(-3, -1)
count = valid.sum(-1, keepdims=True)
_, _, width = _periodic_voronoi_widths(alpha, valid, period)
weight = jnp.where(count == 1, 1.0, width / period)
return jnp.where(valid, weight, 0.0).swapaxes(-3, -1)


def _reduction_gamma_delta(v_tau, radial, poloidal, opts, alpha, mask):
v_tau = (v_tau * _alpha_weights(alpha, mask)).sum(-3)
outward_superbanana = (radial > opts.thresh * jnp.abs(poloidal)).any(-3)
return (v_tau * outward_superbanana).sum(-1)


def _reduction_gamma_alpha(v_tau, radial, poloidal, opts, order=1):
def _reduction_gamma_alpha(v_tau, radial, poloidal, opts, alpha, mask, order=1):
thresh = opts.thresh * jnp.abs(poloidal)
outward_score = radial - thresh
inward_score = -radial - thresh
outward_score = radial - thresh # alpha out candidate
inward_score = -radial - thresh # alpha in candidate

# dist[i,j] is the right-handed distance along unit circle from alpha[i] to alpha[j]
dist = (opts.alpha - opts.alpha[:, None]) % (2 * jnp.pi)
da = 2 * jnp.pi / opts.alpha.size
loss_cone = jnp.where(
poloidal >= 0,
_LossCone.indicator(inward_score, outward_score, dist, da, order=order),
_LossCone.indicator(outward_score, inward_score, dist, da, order=order),
_LossCone.indicator_nonuniform(
inward_score, outward_score, alpha, mask, order=order
),
_LossCone.indicator_nonuniform(
outward_score, inward_score, alpha, mask, order=order
),
)
has_alpha_out = (outward_score > 0).any(-3, keepdims=True)
has_alpha_in = (inward_score > 0).any(-3, keepdims=True)
loss_cone = (has_alpha_out & has_alpha_in) * loss_cone + (
has_alpha_out & ~has_alpha_in
)
return (v_tau * loss_cone).sum(-1).mean(-2)
return (v_tau * loss_cone * _alpha_weights(alpha, mask)).sum((-3, -1))


@register_compute_fun(
Expand Down Expand Up @@ -357,7 +432,13 @@ def _Gamma_delta(params, transforms, profiles, data, **kwargs):
"""Equation 22 of [2]_."""
# noqa: unused dependency
data["Gamma_delta"] = _Gamma(
_reduction_gamma_delta, params, transforms, profiles, data, **kwargs
_reduction_gamma_delta,
params,
transforms,
profiles,
data,
fold_alpha=True,
**kwargs,
)
return data

Expand Down Expand Up @@ -400,32 +481,43 @@ def _Gamma_alpha(params, transforms, profiles, data, **kwargs):
"""Equation 25 of [2]_."""
# noqa: unused dependency
data["Gamma_alpha"] = _Gamma(
_reduction_gamma_alpha, params, transforms, profiles, data, **kwargs
_reduction_gamma_alpha,
params,
transforms,
profiles,
data,
fold_alpha=True,
**kwargs,
)
return data


def _Gamma(reduction, params, transforms, profiles, data, **kwargs):
def _Gamma(reduction, params, transforms, profiles, data, fold_alpha=False, **kwargs):
grid = transforms["grid"]
opts = Options.guess(-1, grid, **kwargs)

def foreach_surface(data):

def foreach(pitch_inv):
return reduction(
*bounce.integrate(
[
_v_tau,
_radial_drift_wb_inverse,
_alpha_drift_wb_inverse,
],
pitch_inv,
data,
names,
num_well=opts.num_well,
),
opts,
points = bounce.points(pitch_inv, opts.num_well) if fold_alpha else None
integrals = bounce.integrate(
[
_v_tau,
_radial_drift_wb_inverse,
_alpha_drift_wb_inverse,
],
pitch_inv,
data,
names,
points,
num_well=opts.num_well,
)
if fold_alpha:
*integrals, alpha, mask = _fold_wells_to_alpha(
integrals, points, opts, data["iota"], grid.NFP
)
return reduction(*integrals, opts, alpha=alpha, mask=mask)
return reduction(*integrals, opts)

pitch_inv, weight = Bounce2D.pitch_quad(
data["min_tz |B|"], data["max_tz |B|"], opts.pitch_quad
Expand All @@ -449,4 +541,6 @@ def foreach(pitch_inv):
)
assert out.ndim == 1
scalar = jnp.pi**3 / 16 * grid.NFP / opts.num_field_periods
if fold_alpha:
scalar *= opts.num_field_periods
return grid.expand(out * scalar) / data["V_psi"]
123 changes: 115 additions & 8 deletions desc/integrals/quad_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -319,7 +319,19 @@ def get_quadrature(quad, automorphism):
return x, w


# This can be made more effecient but it gets the job done.
def _periodic_voronoi_widths(alpha, valid, period=2 * jnp.pi):
"""Periodic Voronoi neighbor distances and cell widths."""
dist = (alpha[..., None, :] - alpha[..., :, None]) % period
dist = jnp.where(
valid[..., :, None] & valid[..., None, :] & (dist > 0), dist, jnp.inf
)
has_neighbors = valid & (valid.sum(-1, keepdims=True) > 1)
prev_width = jnp.where(has_neighbors, dist.min(-2), period)
next_width = jnp.where(has_neighbors, dist.min(-1), period)
width = 0.5 * (prev_width + next_width)
return prev_width, next_width, width


class _LossCone:
"""Utilities for periodic loss-cone indicators."""

Expand All @@ -337,14 +349,109 @@ def _cell_weight(center, stop, dx, period=2 * jnp.pi):
cell_start = center - dx / 2
cell_stop = center + dx / 2
shift = period * jnp.arange(-1, 2)
coverage = jnp.clip(
jnp.minimum(cell_stop[..., None] + shift, stop[..., None])
- jnp.maximum(cell_start[..., None] + shift, 0.0),
0.0,
dx,
).sum(-1)
coverage = (
(
jnp.minimum(cell_stop[..., None] + shift, stop[..., None])
- jnp.maximum(cell_start[..., None] + shift, 0.0)
)
.clip(0.0, dx)
.sum(-1)
)
return coverage / dx

@staticmethod
def _root_nonuniform(score, alpha, valid, period=2 * jnp.pi):
"""Find negative-to-positive crossings on a nonuniform periodic grid."""
# dist[..., i, j] is the forward distance from sample i to sample j.
dist = (alpha[..., None, :] - alpha[..., :, None]) % period
dist = jnp.where(
valid[..., :, None] & valid[..., None, :] & (dist > 0), dist, jnp.inf
)
prev_idx = dist.argmin(axis=-2)
prev_dist = jnp.take_along_axis(dist, prev_idx[..., None, :], axis=-2)[
..., 0, :
]
previous = jnp.take_along_axis(score, prev_idx, axis=-1)
has_previous = jnp.isfinite(prev_dist)
event = valid & has_previous & (score > 0) & (previous <= 0)
prev_dist = jnp.where(has_previous, prev_dist, 0.0)
offset = jnp.where(
has_previous, safediv(prev_dist * score, score - previous), 0.0
)
return event, (alpha - offset) % period

@staticmethod
def _cell_weight_nonuniform(root, stop, alpha, valid, period=2 * jnp.pi):
"""Fraction of each nonuniform periodic cell covered by an interval."""
prev_width, _, width = _periodic_voronoi_widths(alpha, valid, period)
cell_left = alpha - 0.5 * prev_width
left = (cell_left[..., None, :] - root[..., :, None]) % period
right = left + width[..., None, :]
shift = period * jnp.arange(-1, 2)
coverage = (
(
jnp.minimum(right[..., None] + shift, stop[..., None])
- jnp.maximum(left[..., None] + shift, 0.0)
)
.clip(0.0, width[..., None, :, None])
.sum(-1)
)
return coverage / width[..., None, :]

@staticmethod
def indicator_nonuniform(
start_score, stop_score, alpha, valid, period=2 * jnp.pi, order=1
):
"""Periodic interval indicator on a nonuniform alpha grid.

The alpha/sample axis is ``-3`` on input and restored on output.
``order=0`` returns a sampled boolean indicator. ``order=1`` uses
linearly interpolated zero crossings of the signed scores to return
fractional cell weights in ``[0,1]``.

"""
start_score = start_score.swapaxes(-3, -1)
stop_score = stop_score.swapaxes(-3, -1)
alpha = alpha.swapaxes(-3, -1)
valid = valid.swapaxes(-3, -1)

dist = (alpha[..., None, :] - alpha[..., :, None]) % period
if order == 0:
start_sample = (start_score > 0) & valid
stop_sample = (stop_score > 0) & valid
first_stop = jnp.where(stop_sample[..., None, :], dist, jnp.inf).min(
-1, keepdims=True
)
loss_cone = (
start_sample[..., None]
& jnp.isfinite(first_stop)
& valid[..., None, :]
& (dist <= first_stop)
)
return loss_cone.any(-2).swapaxes(-3, -1)

errorif(order != 1, msg="Loss cone indicator order must be 0 or 1.")
start_crossing, start_alpha = _LossCone._root_nonuniform(
start_score, alpha, valid, period
)
stop_crossing, stop_alpha = _LossCone._root_nonuniform(
stop_score, alpha, valid, period
)

stop_dist = (stop_alpha[..., None, :] - start_alpha[..., :, None]) % period
first_stop = jnp.where(stop_crossing[..., None, :], stop_dist, jnp.inf).min(
-1, keepdims=True
)
loss_cone = (
start_crossing[..., None]
* jnp.isfinite(first_stop)
* valid[..., None, :]
* _LossCone._cell_weight_nonuniform(
start_alpha, first_stop, alpha, valid, period
)
)
return loss_cone.sum(-2).clip(0.0, 1.0).swapaxes(-3, -1)

@staticmethod
def indicator(start_score, stop_score, dist, dx=None, period=2 * jnp.pi, order=1):
"""Periodic interval indicator from branch start and stop scores.
Expand Down Expand Up @@ -406,4 +513,4 @@ def indicator(start_score, stop_score, dist, dx=None, period=2 * jnp.pi, order=1
* jnp.isfinite(first_stop)
* _LossCone._cell_weight(center, first_stop, dx, period)
)
return jnp.clip(loss_cone.sum(-2), 0.0, 1.0).swapaxes(-3, -1)
return loss_cone.sum(-2).clip(0.0, 1.0).swapaxes(-3, -1)
Binary file modified tests/inputs/master_compute_data_rpz.pkl
Binary file not shown.
Loading