Skip to content
Merged
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
5 changes: 4 additions & 1 deletion essos/losses.py
Original file line number Diff line number Diff line change
Expand Up @@ -139,7 +139,10 @@ def grad_pytree(self, dofs_pytree) -> dict:
else:
args = tuple(dofs_pytree)
gradient = jax_grad(self.fun, argnums=tuple(range(len(args))))(*args, **self.kwargs)
buffer = self.dependencies_buffer.copy()
# Build a fresh zeros structure locally instead of using the cached
# dependencies_buffer property, which would cache a traced value and
# leak it out of this jit scope (UnexpectedTracerError).
buffer = tree_util.tree_map(jnp.zeros_like, self.dependencies)
for dep, g in zip(self.args_names, gradient):
buffer[dep] = g
return buffer
Expand Down
148 changes: 75 additions & 73 deletions tests/test_multiobjectives.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,31 +4,32 @@
from essos.multiobjectiveoptimizer import MultiObjectiveOptimizer
from essos.coils import Coils,Curves
from essos.fields import BiotSavart
from essos.objective_functions import loss_bdotn_over_b, loss_coil_length, loss_coil_curvature, loss_normB_axis

# test_multiobjectiveoptimizer.py

import jax.numpy as jnp




def surface():
surface.nphi=3
surface.ntheta=3
surface.gamma = jnp.ones((3, 3, 3))
surface.unitnormal = jnp.ones((3, 3, 3))
return surface

def mock_vmec():
vmec = MagicMock()
vmec.nfp = 2
vmec.r_axis = 10.0
vmec.surface = surface()
return vmec



from essos.objective_functions import loss_coil_length, loss_coil_curvature
from essos.surfaces import BdotN_over_B

# test_multiobjectiveoptimizer.py

import jax.numpy as jnp




def surface():
surface.nphi=3
surface.ntheta=3
surface.gamma = jnp.ones((3, 3, 3))
surface.unitnormal = jnp.ones((3, 3, 3))
return surface

def mock_vmec():
vmec = MagicMock()
vmec.nfp = 2
vmec.r_axis = 10.0
vmec.surface = surface()
return vmec



def dummy_loss_fn(field=None, coils=None, vmec=None, surface=None, x=None):
return jnp.sum(x)

Expand Down Expand Up @@ -68,54 +69,55 @@ def loss_fn(curve_dofs, current):
assert jnp.array_equal(gradient_tuple["unused"], gradient["unused"])


@pytest.mark.xfail(reason='test_build_available_inputs uses the old optimizer loss API (x, dofs_curves=, currents_scale=); BdotN_over_B now lives in essos.surfaces with signature (surface, field). Needs rewrite to new API.', strict=False)
def test_build_available_inputs( vmec=mock_vmec(), dummy_loss_fn=dummy_loss_fn):
Comment thread
EstevaoMGomes marked this conversation as resolved.
optimizer = MultiObjectiveOptimizer(
loss_functions=[dummy_loss_fn],
vmec=vmec,
coils_init=None,
function_inputs={"extra": 42},
opt_config={"order_Fourier": 2, "num_coils": 2}
)
x = jnp.arange(32, dtype=float)
result = optimizer._build_available_inputs(x)
expected_keys = {
"field", "coils", "vmec", "surface", "x", "dofs_curves", "currents_scale", "nfp", "extra"
}
assert expected_keys.issubset(result.keys())
assert isinstance(result["x"], jnp.ndarray)
assert result["vmec"] is vmec
assert result["surface"] is vmec.surface
assert result["currents_scale"] == 1.0
assert result["nfp"] == 2
assert result["extra"] == 42
assert result["dofs_curves"].shape == (2, 3,5)
weights=jnp.array([1.0])
loss_result=optimizer._call_loss_fn(dummy_loss_fn,result)
assert loss_result.shape == ()
assert loss_result == 496
loss_weight_result=optimizer.weighted_loss( x, weights)
assert loss_weight_result.shape == ()
assert loss_weight_result == 496
optimized_coils=optimizer.optimize_with_optax(weights, method="adam", lr=1e-2)
assert optimized_coils.currents_scale==0.01999998979999997872
dofs_curves=optimized_coils.dofs_curves
currents_scale=optimized_coils.currents_scale
nfp=optimized_coils.nfp
n_segments=optimized_coils.n_segments
stellsym=optimized_coils.stellsym
x=optimized_coils.x
bdotn_b=loss_bdotn_over_b(x,vmec=vmec,dofs_curves=dofs_curves, currents_scale=currents_scale, nfp=nfp, n_segments=n_segments, stellsym=stellsym)
#assert bdotn_b==0.0000000000000037761977058799732810080238
max_length=loss_coil_length(x,dofs_curves=dofs_curves, currents_scale=currents_scale, nfp=nfp, n_segments=n_segments, stellsym=stellsym)
max_curvature=loss_coil_curvature(x,dofs_curves=dofs_curves, currents_scale=currents_scale, nfp=nfp, n_segments=n_segments, stellsym=stellsym)
normB_axis=loss_normB_axis(x,dofs_curves=dofs_curves, currents_scale=currents_scale, nfp=nfp, n_segments=n_segments, stellsym=stellsym)
optimizer.run()
vmec=vmec,
coils_init=None,
function_inputs={"extra": 42},
opt_config={"order_Fourier": 2, "num_coils": 2}
)
x = jnp.arange(32, dtype=float)

result = optimizer._build_available_inputs(x)


expected_keys = {
"field", "coils", "vmec", "surface", "x", "dofs_curves", "currents_scale", "nfp", "extra"
}
assert expected_keys.issubset(result.keys())
assert isinstance(result["x"], jnp.ndarray)
assert result["vmec"] is vmec
assert result["surface"] is vmec.surface
assert result["currents_scale"] == 1.0
assert result["nfp"] == 2
assert result["extra"] == 42
assert result["dofs_curves"].shape == (2, 3,5)

weights=jnp.array([1.0])
loss_result=optimizer._call_loss_fn(dummy_loss_fn,result)
assert loss_result.shape == ()
assert loss_result == 496
loss_weight_result=optimizer.weighted_loss( x, weights)
assert loss_weight_result.shape == ()
assert loss_weight_result == 496

optimized_coils=optimizer.optimize_with_optax(weights, method="adam", lr=1e-2)
assert optimized_coils.currents_scale==0.01999998979999997872

dofs_curves=optimized_coils.dofs_curves
currents_scale=optimized_coils.currents_scale
nfp=optimized_coils.nfp
n_segments=optimized_coils.n_segments
stellsym=optimized_coils.stellsym
x=optimized_coils.x
bdotn_b=loss_bdotn_over_b(x,vmec=vmec,dofs_curves=dofs_curves, currents_scale=currents_scale, nfp=nfp, n_segments=n_segments, stellsym=stellsym)
#assert bdotn_b==0.0000000000000037761977058799732810080238

max_length=loss_coil_length(x,dofs_curves=dofs_curves, currents_scale=currents_scale, nfp=nfp, n_segments=n_segments, stellsym=stellsym)
max_curvature=loss_coil_curvature(x,dofs_curves=dofs_curves, currents_scale=currents_scale, nfp=nfp, n_segments=n_segments, stellsym=stellsym)
normB_axis=loss_normB_axis(x,dofs_curves=dofs_curves, currents_scale=currents_scale, nfp=nfp, n_segments=n_segments, stellsym=stellsym)

optimizer.run()

76 changes: 50 additions & 26 deletions tests/test_objective_functions.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,18 +8,30 @@
import essos.objective_functions as objf


class DummyField:
class DummyCoils:
def __init__(self):
self.R0 = jnp.array([1.0, 1.0])
self.Z0 = jnp.array([0.0, 0.0])
self.phi = jnp.array([0.0, jnp.pi / 2])
self.B_axis = jnp.array([[1.0, 0.0, 0.0], [0.0, 1.0, 0.0]])
self.grad_B_axis = jnp.ones((2, 3, 3))
self.iota = 0.4
self.elongation = jnp.array([1.0, 1.2])
self.r_axis = 1.0
self.z_axis = 0.0
self.coils = DummyCoils()
self.length = jnp.array([3.0, 4.0])
self.curvature = jnp.ones((2, 5))

def copy(self):
return DummyCoils()


@jax.tree_util.register_pytree_node_class
class DummyField:
def __init__(self, R0=None, Z0=None, phi=None, B_axis=None, grad_B_axis=None,
iota=0.4, elongation=None, r_axis=1.0, z_axis=0.0, coils=None):
self.R0 = jnp.array([1.0, 1.0]) if R0 is None else R0
self.Z0 = jnp.array([0.0, 0.0]) if Z0 is None else Z0
self.phi = jnp.array([0.0, jnp.pi / 2]) if phi is None else phi
# B_axis matches pyqsc_jax external module shape (3, n); function applies .T
self.B_axis = jnp.array([[1.0, 0.0], [0.0, 1.0], [0.0, 0.0]]) if B_axis is None else B_axis
self.grad_B_axis = jnp.ones((3, 3, 2)) if grad_B_axis is None else grad_B_axis
self.iota = iota
self.elongation = jnp.array([1.0, 1.2]) if elongation is None else elongation
self.r_axis = r_axis
self.z_axis = z_axis
self.coils = DummyCoils() if coils is None else coils

def AbsB(self, points):
return 5.7 + 0.1 * jnp.sum(points)
Expand All @@ -36,6 +48,17 @@ def B_covariant(self, points):
def copy(self):
return DummyField()

def tree_flatten(self):
children = (self.R0, self.Z0, self.phi, self.B_axis, self.grad_B_axis,
self.iota, self.elongation, self.r_axis, self.z_axis)
return children, {}

@classmethod
def tree_unflatten(cls, aux, children):
(R0, Z0, phi, B_axis, grad_B_axis, iota, elongation, r_axis, z_axis) = children
return cls(R0=R0, Z0=Z0, phi=phi, B_axis=B_axis, grad_B_axis=grad_B_axis,
iota=iota, elongation=elongation, r_axis=r_axis, z_axis=z_axis)


class DummyParticles:
def __init__(self):
Expand Down Expand Up @@ -64,12 +87,22 @@ def __init__(self, *args, **kwargs):
self.maxtime = 1e-5


@jax.tree_util.register_pytree_node_class
class DummySurface:
def __init__(self):
self.gamma = jnp.zeros((2, 3, 3), dtype=jnp.float64)
self.unitnormal = jnp.ones((2, 3, 3), dtype=jnp.float64)
self.stellsym = False
self.nfp = 1
def __init__(self, gamma=None, unitnormal=None, stellsym=False, nfp=1):
self.gamma = jnp.zeros((2, 3, 3), dtype=jnp.float64) if gamma is None else gamma
self.unitnormal = jnp.ones((2, 3, 3), dtype=jnp.float64) if unitnormal is None else unitnormal
self.stellsym = stellsym
self.nfp = nfp

def tree_flatten(self):
return (self.gamma, self.unitnormal), (self.stellsym, self.nfp)

@classmethod
def tree_unflatten(cls, aux, children):
gamma, unitnormal = children
stellsym, nfp = aux
return cls(gamma=gamma, unitnormal=unitnormal, stellsym=stellsym, nfp=nfp)


@jax.tree_util.register_pytree_node_class
Expand Down Expand Up @@ -168,7 +201,7 @@ def test_near_axis_losses(self):
points, B_nearaxis, gradB_nearaxis = objf.near_axis_field_quantities(self.field_nearaxis)
self.assertEqual(points.shape, (3, 2))
self.assertEqual(B_nearaxis.shape, (2, 3))
self.assertEqual(gradB_nearaxis.shape, (3, 3, 2))
self.assertEqual(gradB_nearaxis.shape, (2, 3, 3))
self.assertTrue(jnp.isfinite(objf.loss_B_difference_coils_near_axis(self.field, self.field_nearaxis)))
self.assertTrue(jnp.isfinite(objf.loss_gradB_difference_coils_near_axis(self.field, self.field_nearaxis)))
self.assertTrue(jnp.isfinite(objf.loss_iota_near_axis(self.field_nearaxis)))
Expand Down Expand Up @@ -262,14 +295,5 @@ def test_regularization_helpers(self):
self.assertTrue(jnp.isfinite(objf.rectangular_xsection_delta(2.0, 1.0)))


class DummyCoils:
def __init__(self):
self.length = jnp.array([3.0, 4.0])
self.curvature = jnp.ones((2, 5))

def copy(self):
return DummyCoils()


if __name__ == "__main__":
unittest.main()