diff --git a/debug_from_simsopt_import.py b/debug_from_simsopt_import.py new file mode 100644 index 0000000..bff6816 --- /dev/null +++ b/debug_from_simsopt_import.py @@ -0,0 +1,96 @@ +import os +import jax +import jax.numpy as jnp + +from essos.coils import Coils +from essos.fields import BiotSavart, Vmec +from essos.constants import ( + ALPHA_PARTICLE_MASS, + ALPHA_PARTICLE_CHARGE, + FUSION_ALPHA_PARTICLE_ENERGY, +) +from essos.dynamics import GuidingCenter, Particles +from essos.objective_functions import normB_axis + + +def describe_case(label, coils, vmec_surface, initial_xyz, initial_vpar): + print(f"\n=== {label} ===") + print("n_base_curves:", coils.curves.n_base_curves) + print("n_segments:", coils.n_segments) + print("nfp:", coils.nfp, "stellsym:", coils.stellsym) + print("dofs_curves shape:", coils.dofs_curves.shape) + print("dofs_currents_raw shape:", coils.dofs_currents_raw.shape) + print("currents raw min/max:", float(jnp.min(coils.dofs_currents_raw)), float(jnp.max(coils.dofs_currents_raw))) + print("currents scale:", float(coils.currents_scale)) + print("currents normalized min/max:", float(jnp.min(coils.dofs_currents)), float(jnp.max(coils.dofs_currents))) + print("gamma shape:", coils.gamma.shape) + print("gamma finite:", bool(jnp.all(jnp.isfinite(coils.gamma)))) + print("gamma_dash finite:", bool(jnp.all(jnp.isfinite(coils.gamma_dash)))) + print("gamma_dashdash finite:", bool(jnp.all(jnp.isfinite(coils.gamma_dashdash)))) + + field = BiotSavart(coils) + B_axis = normB_axis(field, npoints=200) + print("axis B mean before renorm:", float(jnp.mean(B_axis))) + coils.dofs_currents = coils.dofs_currents * 5.7 / jnp.mean(B_axis) + field = BiotSavart(coils) + print("axis B mean after renorm:", float(jnp.mean(normB_axis(field, npoints=200)))) + + print("surface gamma finite:", bool(jnp.all(jnp.isfinite(vmec_surface.gamma)))) + print("initial xyz:", initial_xyz) + print("field.B(initial):", field.B(initial_xyz)) + print("field.AbsB(initial):", float(field.AbsB(initial_xyz))) + print("field.dAbsB_by_dX(initial):", field.dAbsB_by_dX(initial_xyz)) + print("field.curl_b(initial):", field.curl_b(initial_xyz)) + print("field.kappa(initial):", field.kappa(initial_xyz)) + + finite_checks = { + "B": jnp.all(jnp.isfinite(field.B(initial_xyz))), + "AbsB": jnp.isfinite(field.AbsB(initial_xyz)), + "dAbsB_by_dX": jnp.all(jnp.isfinite(field.dAbsB_by_dX(initial_xyz))), + "curl_b": jnp.all(jnp.isfinite(field.curl_b(initial_xyz))), + "kappa": jnp.all(jnp.isfinite(field.kappa(initial_xyz))), + } + print("finite checks:", {k: bool(v) for k, v in finite_checks.items()}) + + particles = Particles( + initial_xyz=jnp.expand_dims(initial_xyz, 0), + mass=ALPHA_PARTICLE_MASS, + charge=ALPHA_PARTICLE_CHARGE, + energy=FUSION_ALPHA_PARTICLE_ENERGY, + ) + initial_condition = jnp.array([initial_xyz[0], initial_xyz[1], initial_xyz[2], initial_vpar]) + rhs = GuidingCenter(0.0, initial_condition, (field, particles, type("E", (), {"E_covariant": lambda self, x: jnp.zeros(3)})())) + print("GuidingCenter rhs:", rhs) + print("GuidingCenter rhs finite:", bool(jnp.all(jnp.isfinite(rhs)))) + + +def main(): + print(jax.devices()) + + simsopt_json = os.path.join("examples", "input_files", "QH_simple_scaled.json") + wout_file = os.path.join("examples", "input_files", "wout_QH_simple_scaled.nc") + vmec = Vmec(wout_file) + + R0 = 17.0 + initial_xyz = jnp.array([R0, 0.0, 0.0]) + initial_vpar = 0.0 + + coils_true = Coils.from_simsopt(simsopt_json, nfp=4, stellsym=True) + describe_case("from_simsopt stellsym=True", coils_true, vmec.surface, initial_xyz, initial_vpar) + + coils_false = Coils.from_simsopt(simsopt_json, nfp=4, stellsym=False) + describe_case("from_simsopt stellsym=False", coils_false, vmec.surface, initial_xyz, initial_vpar) + + from simsopt import load + + bs = load(simsopt_json) + base_coils = bs.coils[:6] + print("\n=== direct simsopt base coils ===") + print("n base coils:", len(base_coils)) + print("current values:", [float(c.current.get_value()) for c in base_coils]) + print("curve gamma shape:", jnp.asarray(base_coils[0].curve.gamma()).shape) + print("curve gamma first point:", jnp.asarray(base_coils[0].curve.gamma())[0]) + + +if __name__ == "__main__": + main() diff --git a/essos/augmented_lagrangian.py b/essos/augmented_lagrangian.py index 260323b..f68c500 100644 --- a/essos/augmented_lagrangian.py +++ b/essos/augmented_lagrangian.py @@ -21,6 +21,11 @@ class LagrangeMultiplier(NamedTuple): def _multiplier_like(out, multiplier, penalty, omega, eta, sq_grad): + if out is None: + raise ValueError( + "Constraint function returned None during initialization. " + "Constraints used with eq()/ineq() must return a scalar or array." + ) z = jnp.zeros_like(out) return LagrangeMultiplier( value=multiplier + z, diff --git a/essos/coil_perturbation.py b/essos/coil_perturbation.py index 3d8b423..f29a618 100644 --- a/essos/coil_perturbation.py +++ b/essos/coil_perturbation.py @@ -238,7 +238,7 @@ def perturb_curves(curves, sampler:GaussianSampler, key=None, perturbation_type= or "statistical"/"statistic" to perturb every expanded coil independently. Returns: - A new perturbed object of the same family as the input. + A new perturbed object of the same type as the input. """ if perturbation_type == "systematic": if isinstance(curves, DiscretizedCoils): @@ -248,19 +248,21 @@ def perturb_curves(curves, sampler:GaussianSampler, key=None, perturbation_type= currents=curves.dofs_currents_raw, nfp=curves.nfp, stellsym=curves.stellsym, + currents_scale=curves.currents_scale, + scale_fixed=curves.scale_fixed, ) if isinstance(curves, Coils): perturbation = _draw_curve_perturbation(sampler, key, curves.curves.n_base_curves) - base_curves = _make_curves_like(curves.curves, curves.dofs_curves, nfp=1, stellsym=False) + base_curves = _make_curves_like(curves.curves, curves.curves._dofs, nfp=1, stellsym=False) perturbed_base_gamma = base_curves.gamma + perturbation dofs_new, _ = fit_dofs_from_coils(perturbed_base_gamma, curves.order, curves.n_segments, assume_uniform=True) new_curves = _make_curves_like(curves.curves, dofs_new, nfp=curves.nfp, stellsym=curves.stellsym) - return Coils(curves=new_curves, currents=curves.dofs_currents_raw) + return Coils(curves=new_curves, currents=curves.dofs_currents_raw, currents_scale=curves.currents_scale) if isinstance(curves, Curves): perturbation = _draw_curve_perturbation(sampler, key, curves.n_base_curves) - base_curves = _make_curves_like(curves, curves.dofs, nfp=1, stellsym=False) + base_curves = _make_curves_like(curves, curves._dofs, nfp=1, stellsym=False) perturbed_base_gamma = base_curves.gamma + perturbation dofs_new, _ = fit_dofs_from_coils(perturbed_base_gamma, curves.order, curves.n_segments, assume_uniform=True) return _make_curves_like(curves, dofs_new, nfp=curves.nfp, stellsym=curves.stellsym) @@ -270,12 +272,19 @@ def perturb_curves(curves, sampler:GaussianSampler, key=None, perturbation_type= gamma_perturbed = curves.gamma + perturbation if isinstance(curves, DiscretizedCoils): - return DiscretizedCoils(gamma_perturbed, currents=curves.currents, nfp=1, stellsym=False) + return DiscretizedCoils( + gamma_perturbed, + currents=curves.currents, + nfp=1, + stellsym=False, + currents_scale=curves.currents_scale, + scale_fixed=curves.scale_fixed, + ) if isinstance(curves, Coils): dofs_new, _ = fit_dofs_from_coils(gamma_perturbed, curves.order, curves.n_segments, assume_uniform=True) new_curves = _make_curves_like(curves.curves, dofs_new, nfp=1, stellsym=False) - return Coils(curves=new_curves, currents=curves.currents) + return Coils(curves=new_curves, currents=curves.currents, currents_scale=curves.currents_scale) if isinstance(curves, Curves): dofs_new, _ = fit_dofs_from_coils(gamma_perturbed, curves.order, curves.n_segments, assume_uniform=True) diff --git a/essos/coils.py b/essos/coils.py index c4c4c3c..ec1cba1 100644 --- a/essos/coils.py +++ b/essos/coils.py @@ -58,17 +58,26 @@ def __init__(self, assert isinstance(nfp, int), "nfp must be a positive integer" assert nfp > 0, "nfp must be a positive integer" assert isinstance(stellsym, bool), "stellsym must be a boolean" - + self._initialize_state( + dofs, + n_segments, + nfp, + stellsym, + self._normalize_scaling_type(scaling_type), + scaling_factor, + scale_fixed, + ) + + def _initialize_state(self, dofs, n_segments, nfp, stellsym, scaling_type, scaling_factor, scale_fixed, order=None): self._dofs = dofs self._n_segments = n_segments self._nfp = nfp self._stellsym = stellsym - - self._scaling_type = self._normalize_scaling_type(scaling_type) + self._scaling_type = scaling_type self._scaling_factor = scaling_factor self._scale_fixed = scale_fixed self._scaling = None - + self._order = dofs.shape[2] // 2 if hasattr(dofs, "shape") else order self.quadpoints = jnp.linspace(0, 1, self._n_segments, endpoint=False) self._curves = None self._gamma = None @@ -119,6 +128,7 @@ def dofs(self): def dofs(self, new_dofs): self.reset_cache() self._dofs = new_dofs / self.scaling[None, None, :] + self._order = self._dofs.shape[2] // 2 # n_segments property and setter @property @@ -186,29 +196,34 @@ def scale_fixed(self, new_scale): def scaling(self): """Mode-by-mode scaling ``scale_fixed * exp(scaling_factor * ||mode_orders||)``.""" if self._scaling is None: - self._scaling = self._compute_mode_scaling( + scaling = self._compute_mode_scaling( self.order, self.scaling_type, self.scaling_factor, self.scale_fixed ) + if not isinstance(scaling, jax.core.Tracer): + self._scaling = scaling + return scaling return self._scaling # order property and setter @property def order(self): - return self._dofs.shape[2]//2 + if hasattr(self._dofs, "shape"): + return self._dofs.shape[2] // 2 + return self._order @order.setter def order(self, new_order): self.reset_cache() # Get unscaled dofs, resize, then store unscaled - old_scaling = self.scaling unscaled_dofs = self._dofs self._dofs = jnp.pad(unscaled_dofs, ((0,0), (0,0), (0, max(0, 2*(new_order-self.order)))))[:, :, :2*(new_order)+1] self._scaling = None # Force recalculation for new order + self._order = new_order # n_base_curves property @property def n_base_curves(self): - return self.dofs.shape[0] + return self._dofs.shape[0] # curves property @property @@ -288,8 +303,18 @@ def curvature(self): # copy method def copy(self): - deep_copy = tree_util.tree_map(lambda x: x.copy(), self) - return deep_copy + curves = object.__new__(Curves) + curves._initialize_state( + self._dofs.copy(), + self._n_segments, + self._nfp, + self._stellsym, + self._scaling_type, + self._scaling_factor, + self._scale_fixed, + order=self.order, + ) + return curves # magic methods def __str__(self): @@ -448,31 +473,48 @@ def from_simsopt(cls, simsopt_curves, nfp=1, stellsym=True, scaling_type=2, scal simsopt_coils = bs.coils simsopt_curves = [c.curve for c in simsopt_coils] simsopt_curves = simsopt_curves[0:int(len(simsopt_curves)/nfp/(1+stellsym))] - dofs = jnp.reshape(jnp.array( - [curve.x for curve in simsopt_curves] + dofs = jnp.reshape(jnp.asarray( + [jnp.asarray(curve.x, dtype=float) for curve in simsopt_curves], + dtype=float, ), (len(simsopt_curves), 3, 2*simsopt_curves[0].order+1)) n_segments = len(simsopt_curves[0].quadpoints) return cls(dofs, n_segments, nfp, stellsym, scaling_type, scaling_factor, scale_fixed) def _tree_flatten(self): - children = (self.dofs,) # arrays / dynamic values + dofs = self.dofs if hasattr(self._dofs, "shape") else self._dofs + children = (dofs,) # arrays / dynamic values aux_data = {"n_segments": self._n_segments, "nfp": self._nfp, "stellsym": self._stellsym, "scaling_type": self._scaling_type, "scaling_factor": self._scaling_factor, - "scale_fixed": self._scale_fixed} # static values + "scale_fixed": self._scale_fixed, + "order": self.order} # static values return (children, aux_data) @classmethod def _tree_unflatten(cls, aux_data, children): dofs, = children - order = dofs.shape[2] // 2 - scaling_type = cls._normalize_scaling_type(aux_data["scaling_type"]) - scaling = cls._compute_mode_scaling( - order, scaling_type, aux_data["scaling_factor"], aux_data["scale_fixed"] + if hasattr(dofs, "shape"): + scaling = cls._compute_mode_scaling( + aux_data["order"], + aux_data["scaling_type"], + aux_data["scaling_factor"], + aux_data["scale_fixed"], + ) + dofs = dofs / scaling[None, None, :] + obj = object.__new__(cls) + obj._initialize_state( + dofs, + aux_data["n_segments"], + aux_data["nfp"], + aux_data["stellsym"], + aux_data["scaling_type"], + aux_data["scaling_factor"], + aux_data["scale_fixed"], + order=aux_data["order"], ) - return cls(dofs / scaling[None, None, :], **aux_data) + return obj tree_util.register_pytree_node(Curves, Curves._tree_flatten, @@ -481,10 +523,28 @@ def _tree_unflatten(cls, aux_data, children): def _initialize_currents_scale(currents, currents_scale): """Return a fixed current scale for normalized current dofs.""" + currents = jnp.atleast_1d(jnp.asarray(currents)) if currents_scale is None: return jnp.mean(jnp.abs(currents)) return currents_scale +def _normalize_base_currents(currents, curves): + """Return base currents as a 1D array matching the number of base curves.""" + currents = jnp.atleast_1d(jnp.asarray(currents)) + if hasattr(curves._dofs, "shape"): + n_base_curves = curves._dofs.shape[0] + if currents.shape[0] == 1 and n_base_curves != 1: + currents = jnp.full((n_base_curves,), currents[0]) + return currents + + +def _currents_as_array(currents): + if isinstance(currents, bool): + return None + if isinstance(currents, (list, tuple)) or hasattr(currents, "shape") or jnp.isscalar(currents): + return jnp.atleast_1d(jnp.asarray(currents)) + return None + def _initialize_scale_fixed(gamma, scale_fixed): """Return a fixed geometry scale for normalized gamma dofs.""" @@ -519,11 +579,20 @@ def __init__(self, curves: Curves, currents: jnp.ndarray, currents_scale=None): # if hasattr(curves, 'n_base_curves') and hasattr(currents, 'size'): # assert curves.n_base_curves == currents.size, "Number of base curves and number of currents must be the same" - self.curves = curves - self._dofs_currents_raw = currents # Non-normalized base currents + currents_array = _currents_as_array(currents) + if currents_array is not None: + currents_scale = _initialize_currents_scale(currents_array, currents_scale) + self._initialize_state(curves, currents, currents_scale) - self._currents_scale = _initialize_currents_scale(currents, currents_scale) - self._dofs_currents = None + def _initialize_state(self, curves, currents_raw, currents_scale): + self.curves = curves + currents_array = _currents_as_array(currents_raw) + if currents_array is not None: + currents_raw = currents_array + currents_raw = _normalize_base_currents(currents_raw, curves) + self._dofs_currents_raw = currents_raw + self._currents_scale = currents_scale + self._dofs_currents = None if hasattr(currents_raw, "shape") else currents_raw self._currents = None # reset_cache method @@ -548,7 +617,7 @@ def dofs_currents_raw(self): @dofs_currents_raw.setter def dofs_currents_raw(self, new_dofs_currents_raw): self.reset_cache() - self._dofs_currents_raw = new_dofs_currents_raw + self._dofs_currents_raw = jnp.atleast_1d(jnp.asarray(new_dofs_currents_raw)) # currents_scale property and setter @property @@ -564,21 +633,20 @@ def currents_scale(self, new_currents_scale): # dofs_currents property and setter @property def dofs_currents(self): + # Sentinel leaf during PyTree traversal: pass through, don't scale. + if self._dofs_currents_raw is None or isinstance(self._dofs_currents_raw, bool): + return self._dofs_currents_raw if self._dofs_currents is None: - self._dofs_currents = self.dofs_currents_raw / self.currents_scale + dofs_currents = self.dofs_currents_raw / self.currents_scale + if not isinstance(dofs_currents, jax.core.Tracer): + self._dofs_currents = dofs_currents + return dofs_currents return self._dofs_currents @dofs_currents.setter def dofs_currents(self, new_dofs_currents): self.dofs_currents_raw = new_dofs_currents * self.currents_scale - # currents property - @property - def currents(self): - if self._currents is None: - self._currents = apply_symmetries_to_currents(self.dofs_currents_raw, self.nfp, self.stellsym) - return self._currents - # dofs property and setter @property def dofs(self): @@ -604,7 +672,7 @@ def x(self, new_dofs): @property def currents(self): if self._currents is None: - self._currents = apply_symmetries_to_currents(self.dofs_currents*self.currents_scale, self.nfp, self.stellsym) + self._currents = apply_symmetries_to_currents(self.dofs_currents_raw, self.nfp, self.stellsym) return self._currents # gamma property @@ -658,10 +726,10 @@ def n_segments(self, new_n_segments): # copy method def copy(self): - coils = Coils(self.curves.copy(), self.dofs_currents_raw.copy(), currents_scale=self.currents_scale) + coils = Coils(self.curves.copy(), self._dofs_currents_raw.copy(), currents_scale=self.currents_scale) # Initialize caches - coils._dofs_currents = self.dofs_currents + coils._dofs_currents = self._dofs_currents coils._currents = self._currents return coils @@ -796,8 +864,21 @@ def from_simsopt(cls, simsopt_coils, nfp=1, stellsym=True, scaling_type=2, scali bs = load(simsopt_coils) simsopt_coils = bs.coils curves = [c.curve for c in simsopt_coils] - currents = jnp.array([c.current.get_value() for c in simsopt_coils[0:int(len(simsopt_coils)/nfp/(1+stellsym))]]) - return cls(Curves.from_simsopt(curves, nfp, stellsym, scaling_type, scaling_factor, scale_fixed), currents) + curves_obj = Curves.from_simsopt(curves, nfp, stellsym, scaling_type, scaling_factor, scale_fixed) + curves = Curves( + jnp.asarray(curves_obj._dofs, dtype=float), + curves_obj.n_segments, + curves_obj.nfp, + curves_obj.stellsym, + curves_obj.scaling_type, + curves_obj.scaling_factor, + curves_obj.scale_fixed, + ) + currents = jnp.asarray( + [float(c.current.get_value()) for c in simsopt_coils[0:int(len(simsopt_coils)/nfp/(1+stellsym))]], + dtype=float, + ) + return cls(curves, currents) @classmethod def from_json(cls, filename: str): @@ -855,7 +936,11 @@ def _tree_flatten(self): @classmethod def _tree_unflatten(cls, aux_data, children): curves, dofs_currents = children - return cls(curves, dofs_currents * aux_data["currents_scale"], currents_scale=aux_data["currents_scale"]) + if hasattr(dofs_currents, "shape"): + dofs_currents = dofs_currents * aux_data["currents_scale"] + obj = object.__new__(cls) + obj._initialize_state(curves, dofs_currents, aux_data["currents_scale"]) + return obj tree_util.register_pytree_node(Coils, Coils._tree_flatten, @@ -935,11 +1020,12 @@ def apply_symmetries_to_gammas(base_gammas, nfp, stellsym): @partial(jit, static_argnames=['nfp', 'stellsym']) def apply_symmetries_to_currents(base_currents, nfp, stellsym): + base_currents = jnp.atleast_1d(jnp.asarray(base_currents)) flip_list = [False, True] if stellsym else [False] currents = [] for k in range(0, nfp): for flip in flip_list: - for i in range(len(base_currents)): + for i in range(base_currents.shape[0]): current = -base_currents[i] if flip else base_currents[i] currents.append(current) return jnp.array(currents) diff --git a/essos/objective_functions.py b/essos/objective_functions.py index 086936e..e88f4f2 100644 --- a/essos/objective_functions.py +++ b/essos/objective_functions.py @@ -10,7 +10,7 @@ from essos.surfaces import BdotN_over_B from essos.coils import Curves, Coils from essos.constants import mu_0 -from essos.coil_perturbation import perturb_curves +from essos.coil_perturbation import perturb_curves, perturb_curves_systematic, perturb_curves_statistic @@ -203,7 +203,7 @@ def perturbed_field_from_field(field, key, sampler): coils = copy_coils_from_field(field) base_key = jax.random.key(key) split_keys = jax.random.split(base_key, 2) - perturb_curves_systematic(coils, sampler, key=split_keys[0]) + coils = perturb_curves_systematic(coils, sampler, key=split_keys[0]) coils = perturb_curves_statistic(coils, sampler, key=split_keys[1]) return BiotSavart(coils) diff --git a/essos/surfaces.py b/essos/surfaces.py index 63df5ff..3dd3146 100644 --- a/essos/surfaces.py +++ b/essos/surfaces.py @@ -138,7 +138,21 @@ def __init__(self, rc, zs, nfp, mpol, ntor, ntheta=30, nphi=30, close=True, rang assert isinstance(nphi, int) and nphi > 0, "nphi must be a positive integer." assert isinstance(close, bool), "close must be a boolean." assert range_torus in ['full torus', 'half period'], f"Unknown range_torus: {range_torus}. Choose 'full torus' or 'half period'." + self._initialize_state( + rc, + zs, + nfp, + mpol, + ntor, + ntheta, + nphi, + close, + range_torus, + self._normalize_scaling_type(scaling_type), + scaling_factor, + ) + def _initialize_state(self, rc, zs, nfp, mpol, ntor, ntheta, nphi, close, range_torus, scaling_type, scaling_factor): self._rc = rc self._zs = zs self._nfp = nfp @@ -164,8 +178,7 @@ def __init__(self, rc, zs, nfp, mpol, ntor, ntheta=30, nphi=30, close=True, rang self._theta2d = None self._phi2d = None self._angles = None - - self._scaling_type = self._normalize_scaling_type(scaling_type) + self._scaling_type = scaling_type self._scaling_factor = scaling_factor self._scaling = None @@ -183,6 +196,10 @@ def _normalize_scaling_type(scaling_type): "Expected 'L1', 1, 'L2', 2, 'Linfty', -1, or jnp.inf." ) + @staticmethod + def _compute_scaling(xm, xn, scaling_type, scaling_factor): + return jnp.exp(scaling_factor * jnp.linalg.norm(jnp.vstack([xm, xn]), ord=scaling_type, axis=0)) + @classmethod def from_input_file(cls, file, ntheta=30, nphi=30, close=True, range_torus='full torus'): @@ -404,7 +421,10 @@ def scaling_factor(self, new_factor): def scaling(self): """Mode-by-mode scaling ``exp(scaling_factor * ||(xm, xn)||)``.""" if self._scaling is None: - self._scaling = jnp.exp(self.scaling_factor * jnp.linalg.norm(jnp.vstack([self.xm, self.xn]), ord=self.scaling_type, axis=0)) + scaling = self._compute_scaling(self.xm, self.xn, self.scaling_type, self.scaling_factor) + if not isinstance(scaling, jax.core.Tracer): + self._scaling = scaling + return scaling return self._scaling # dofs property and setter @@ -689,7 +709,10 @@ def mean_cross_sectional_area(self): return mean_cross_sectional_area def _tree_flatten(self): - children = (self.dofs,) # arrays / dynamic values + if hasattr(self._rc, "shape") and hasattr(self._zs, "shape"): + children = (self.rc * self.scaling, self.zs * self.scaling) # arrays / dynamic values + else: + children = (self._rc, self._zs) aux_data = {"nfp": self._nfp, "mpol": self._mpol, "ntor": self._ntor, @@ -703,24 +726,40 @@ def _tree_flatten(self): @classmethod def _tree_unflatten(cls, aux_data, children): - dofs, = children - half = dofs.size // 2 - rc_scaled = dofs[:half] - zs_scaled = dofs[half:] - - mpol = aux_data["mpol"] - ntor = aux_data["ntor"] - nfp = aux_data["nfp"] - scaling_type = cls._normalize_scaling_type(aux_data["scaling_type"]) - scaling_factor = aux_data["scaling_factor"] - - xm = jnp.repeat(jnp.arange(mpol + 1), 2 * ntor + 1)[ntor:] - xn = nfp * jnp.tile(jnp.arange(-ntor, ntor + 1), mpol + 1)[ntor:] - scaling = jnp.exp(scaling_factor * jnp.linalg.norm(jnp.vstack([xm, xn]), ord=scaling_type, axis=0)) - - rc = rc_scaled / scaling - zs = zs_scaled / scaling - return cls(rc, zs, **aux_data) + rc_scaled, zs_scaled = children + + if hasattr(rc_scaled, "shape") and hasattr(zs_scaled, "shape"): + mpol = aux_data["mpol"] + ntor = aux_data["ntor"] + nfp = aux_data["nfp"] + scaling_type = cls._normalize_scaling_type(aux_data["scaling_type"]) + scaling_factor = aux_data["scaling_factor"] + + xm = jnp.repeat(jnp.arange(mpol + 1), 2 * ntor + 1)[ntor:] + xn = nfp * jnp.tile(jnp.arange(-ntor, ntor + 1), mpol + 1)[ntor:] + scaling = cls._compute_scaling(xm, xn, scaling_type, scaling_factor) + + rc = rc_scaled / scaling + zs = zs_scaled / scaling + else: + rc = rc_scaled + zs = zs_scaled + + obj = object.__new__(cls) + obj._initialize_state( + rc, + zs, + aux_data["nfp"], + aux_data["mpol"], + aux_data["ntor"], + aux_data["ntheta"], + aux_data["nphi"], + aux_data["close"], + aux_data["range_torus"], + aux_data["scaling_type"], + aux_data["scaling_factor"], + ) + return obj tree_util.register_pytree_node(SurfaceRZFourier, SurfaceRZFourier._tree_flatten, diff --git a/examples/coil_optimization/optimize_coils_and_nearaxis.py b/examples/coil_optimization/optimize_coils_and_nearaxis.py index e842f07..8be7ce5 100644 --- a/examples/coil_optimization/optimize_coils_and_nearaxis.py +++ b/examples/coil_optimization/optimize_coils_and_nearaxis.py @@ -1,5 +1,5 @@ import os -number_of_processors_to_use = 4 # Parallelization, this should divide nfieldlines +number_of_processors_to_use = 1 # Parallelization, this should divide nfieldlines os.environ["XLA_FLAGS"] = f'--xla_force_host_platform_device_count={number_of_processors_to_use}' from time import time import jax.numpy as jnp @@ -9,72 +9,141 @@ from pyqsc_jax.near_axis import near_axis from essos.dynamics import Tracing from essos.optimization import optimize_loss_function -from essos.objective_functions import (difference_B_gradB_onaxis, - loss_coils_and_nearaxis, loss_coils_for_nearaxis) - -# Optimization parameters -max_coil_length = 4. -max_coil_curvature = 6. -order_Fourier_series_coils = 5 -number_coil_points = order_Fourier_series_coils*10 -maximum_function_evaluations = 200 -number_coils_per_half_field_period = 3 +from jax import vmap, jit +# In this exmple, `scipy.optimize.least_squares` is used, but any other optimizer, e.g. from +# `scipy.optimize.minimize` or `jaxopt`, can be used as well and may even be preferable. +from scipy.optimize import least_squares +from essos.losses import custom_loss + + +""" Creating starting coils and surface """ +N_COILS = 3; FOURIER_ORDER = 6; LARGE_R = 10; SMALL_R = 5.6; NFP = 3; N_SEGMENTS = 60; STELLSYM = True # Curve parameters +COIL_CURRENT = 1. # Amperes (optimization does not depend on current magnitude) tolerance_optimization = 1e-8 +maximum_function_evaluations = 200 + + # Initialize Near-Axis field -rc=jnp.array([1, 0.045]) -zs=jnp.array([0,-0.045]) +rc=jnp.array([1.,0.01]) +zs=jnp.array([0.,0.01]) etabar=-0.9 -nfp=3 -field_nearaxis_initial = near_axis(rc=rc, zs=zs, etabar=etabar, nfp=nfp) +field_nearaxis_initial = near_axis(rc=rc, zs=zs, etabar=etabar, nfp=NFP,order='r2') # Initialize coils -current_on_each_coil = 17e5*field_nearaxis_initial.B0/nfp/2 -number_of_field_periods = nfp +current_on_each_coil = 17.e5*field_nearaxis_initial.B0/NFP/2. +number_of_field_periods = NFP major_radius_coils = field_nearaxis_initial.R0[0] minor_radius_coils = major_radius_coils/2.0 -curves = CreateEquallySpacedCurves(n_curves=number_coils_per_half_field_period, - order=order_Fourier_series_coils, +init_curves = CreateEquallySpacedCurves(n_curves=N_COILS, + order=FOURIER_ORDER, R=major_radius_coils, r=minor_radius_coils, - n_segments=number_coil_points, - nfp=number_of_field_periods, stellsym=True) -coils_initial = Coils(curves=curves, currents=[current_on_each_coil]*number_coils_per_half_field_period) + n_segments=N_SEGMENTS, + nfp=number_of_field_periods, stellsym=STELLSYM) +init_coils = Coils(curves=init_curves, currents=jnp.array([current_on_each_coil]*N_COILS)) +init_field = BiotSavart(init_coils) -# Optimize coils -print(f'Optimizing coils for initial near=axis with {maximum_function_evaluations} function evaluations.') -time0 = time() -initial_dofs = coils_initial.x -coils_optimized_initial_nearaxis = optimize_loss_function(loss_coils_for_nearaxis, initial_dofs=coils_initial.x, coils=coils_initial, tolerance_optimization=tolerance_optimization, - maximum_function_evaluations=maximum_function_evaluations, field_nearaxis=field_nearaxis_initial, - max_coil_length=max_coil_length, max_coil_curvature=max_coil_curvature,) -print(f"Optimization took {time()-time0:.2f} seconds") - -# Optimize coils -print(f'Optimizing coils and near-axis with {maximum_function_evaluations} function evaluations.') -time0 = time() -initial_dofs = jnp.concatenate((coils_optimized_initial_nearaxis.x, field_nearaxis_initial.x)) -coils_optimized, field_nearaxis_optimized = optimize_loss_function(loss_coils_and_nearaxis, initial_dofs=initial_dofs, coils=coils_initial, tolerance_optimization=tolerance_optimization, - maximum_function_evaluations=maximum_function_evaluations, field_nearaxis=field_nearaxis_initial, - max_coil_length=max_coil_length, max_coil_curvature=max_coil_curvature,) -print(f"Optimization took {time()-time0:.2f} seconds") -B_difference_initial, gradB_difference_initial = difference_B_gradB_onaxis(field_nearaxis_initial, BiotSavart(coils_optimized_initial_nearaxis)) -B_difference_loss_initial = jnp.sum(jnp.abs(B_difference_initial)) -gradB_difference_loss_initial = jnp.sum(jnp.abs(gradB_difference_initial)) +""" Setting the losses weights and targets """ +LENGTH_WEIGHT = 1.; LENGTH_TARGET = 4. +CURVATURE_WEIGHT = 1.; CURVATURE_TARGET = 6. +B_DIFFERENCE_WEIGHT = 1. +GRADB_DIFFERENCE_WEIGHT = 1. +IOTA_TARGET = 0.41 +IOTA_WEIGHT = 10. +R0_TARGET = field_nearaxis_initial.R0[0] +R0_WEIGHT = 10. + + + +""" Creating the loss functions """ +def near_axis_field_quantities(field_nearaxis): + Raxis = field_nearaxis.R0 + Zaxis = field_nearaxis.Z0 + phi = field_nearaxis.phi + Xaxis = Raxis*jnp.cos(phi) + Yaxis = Raxis*jnp.sin(phi) + points = jnp.array([Xaxis, Yaxis, Zaxis]) + B_nearaxis = field_nearaxis.B_axis.T + gradB_nearaxis = field_nearaxis.grad_B_axis.T + return points, B_nearaxis, gradB_nearaxis + + + +def loss_B_difference_coils_near_axis(field, field_nearaxis): + points, B_nearaxis, _ = near_axis_field_quantities(field_nearaxis) + B_coils = vmap(field.B)(points.T) + B_difference_loss = jnp.sum(jnp.abs(jnp.array(B_coils)-jnp.array(B_nearaxis))) + return B_difference_loss + +def loss_gradB_difference_coils_near_axis(field, field_nearaxis): + points, _, gradB_nearaxis = near_axis_field_quantities(field_nearaxis) + gradB_coils = vmap(field.dB_by_dX)(points.T) + gradB_difference_loss = jnp.sum(jnp.abs(jnp.array(gradB_coils)-jnp.array(gradB_nearaxis))) + return gradB_difference_loss + +def loss_iota_near_axis(field_nearaxis,iota_target=IOTA_TARGET): + return jnp.abs((field_nearaxis.iota - iota_target)) + +def loss_r0_near_axis(field_nearaxis, r0_target=R0_TARGET): + return jnp.abs((field_nearaxis.R0[0] - r0_target)) + +def loss_length(field,length_target=LENGTH_TARGET): + return jnp.mean(jnp.maximum(0, field.coils.length - length_target)) + +def loss_curvature(field,curvature_target=CURVATURE_TARGET): + return jnp.mean(jnp.maximum(0, field.coils.curvature - curvature_target)) + + +""" Defining custom losses """ +L_B_difference = custom_loss(loss_B_difference_coils_near_axis, "field", "field_nearaxis") +L_gradB_difference = custom_loss(loss_gradB_difference_coils_near_axis, "field", "field_nearaxis") +L_length = custom_loss(loss_length, "field") +L_curvature = custom_loss(loss_curvature, "field") +L_iota = custom_loss(loss_iota_near_axis, "field_nearaxis") +L_r0 = custom_loss(loss_r0_near_axis, "field_nearaxis") + + +""" Defining total loss + setting dependencies """ +L_total = B_DIFFERENCE_WEIGHT*L_B_difference + GRADB_DIFFERENCE_WEIGHT*L_gradB_difference + LENGTH_WEIGHT*L_length + CURVATURE_WEIGHT*L_curvature + IOTA_WEIGHT*L_iota+ R0_WEIGHT*L_r0 + + +L_total.dependencies = {"field": init_field, "field_nearaxis": field_nearaxis_initial} + +""" Optimizing the total loss """ +t_start = time() +res = least_squares(L_total, L_total.starting_dofs, L_total.grad, verbose=2, ftol=1e-5, gtol=1e-5, xtol=1e-14, max_nfev=maximum_function_evaluations) +t_end = time() + +print(f"\nOptimization took {t_end - t_start:.2f} seconds") +print("Initial loss:", L_total(L_total.starting_dofs)) +print("Loss after optimization:", L_total(res.x)) + +opt_field = L_total.dofs_to_pytree(res.x)["field"] +opt_coils = opt_field.coils + +opt_field_nearaxis = L_total.dofs_to_pytree(res.x)["field_nearaxis"] + + +B_difference_initial = loss_B_difference_coils_near_axis(init_field, field_nearaxis_initial) +gradB_difference_initial = loss_gradB_difference_coils_near_axis(init_field, field_nearaxis_initial) + +B_difference_optimized = loss_B_difference_coils_near_axis(opt_field, opt_field_nearaxis) +gradB_difference_optimized = loss_gradB_difference_coils_near_axis(opt_field, opt_field_nearaxis) -B_difference_optimized, gradB_difference_optimized = difference_B_gradB_onaxis(field_nearaxis_optimized, BiotSavart(coils_optimized)) -B_difference_loss_optimized = jnp.sum(jnp.abs(B_difference_optimized)) -gradB_difference_loss_optimized = jnp.sum(gradB_difference_optimized) print(f'############################################') print(f'Iota for initial near-axis: {field_nearaxis_initial.iota}') -print(f'Iota for optimized near-axis: {field_nearaxis_optimized.iota}') +print(f'Iota for optimized near-axis: {opt_field_nearaxis.iota}') print(f'Maximum elongation for initial near-axis: {max(field_nearaxis_initial.elongation)}') -print(f'Maximum elongation for optimized near-axis: {max(field_nearaxis_optimized.elongation)}') -print(f'Loss of B difference for initial near-axis: {B_difference_loss_initial}') -print(f'Loss of B difference for optimized near-axis: {B_difference_loss_optimized}') -print(f'Loss of gradB difference for initial near-axis: {gradB_difference_loss_initial}') -print(f'Loss of gradB difference for optimized near-axis: {gradB_difference_loss_optimized}') +print(f'Maximum elongation for optimized near-axis: {max(opt_field_nearaxis.elongation)}') +print(f'Loss of B difference for initial near-axis: {B_difference_initial}') +print(f'Loss of B difference for optimized near-axis: {B_difference_optimized}') +print(f'Loss of gradB difference for initial near-axis: {gradB_difference_initial}') +print(f'Loss of gradB difference for optimized near-axis: {gradB_difference_optimized}') +print(f'Loss of R0 difference for initial near-axis: {loss_r0_near_axis(field_nearaxis_initial)}') +print(f'Loss of R0 difference for optimized near-axis: {loss_r0_near_axis(opt_field_nearaxis)}') + # Trace fieldlines nfieldlines = 6 @@ -83,16 +152,16 @@ trace_tolerance = 1e-7 R0_initial = jnp.linspace(field_nearaxis_initial.R0[0], 1.05*field_nearaxis_initial.R0[0], nfieldlines) -R0_optimized = jnp.linspace(field_nearaxis_optimized.R0[0], 1.05*field_nearaxis_optimized.R0[0], nfieldlines) +R0_optimized = jnp.linspace(opt_field_nearaxis.R0[0], 1.05*opt_field_nearaxis.R0[0], nfieldlines) Z0 = jnp.zeros(nfieldlines) phi0 = jnp.zeros(nfieldlines) initial_xyz_initial = jnp.array([R0_initial*jnp.cos(phi0), R0_initial*jnp.sin(phi0), Z0]).T initial_xyz_optimized = jnp.array([R0_optimized*jnp.cos(phi0), R0_optimized*jnp.sin(phi0), Z0]).T time0 = time() -tracing_initial = Tracing(field=BiotSavart(coils_optimized_initial_nearaxis), model='FieldLineAdaptative', initial_conditions=initial_xyz_initial, +tracing_initial = Tracing(field=init_field, model='FieldLineAdaptative', initial_conditions=initial_xyz_initial, maxtime=tmax, times_to_trace=num_steps, atol=trace_tolerance,rtol=trace_tolerance) -tracing_optimized = Tracing(field=BiotSavart(coils_optimized), model='FieldLineAdaptative', initial_conditions=initial_xyz_optimized, +tracing_optimized = Tracing(field=opt_field, model='FieldLineAdaptative', initial_conditions=initial_xyz_optimized, maxtime=tmax, times_to_trace=num_steps, atol=trace_tolerance,rtol=trace_tolerance) print(f"Tracing took {time()-time0:.2f} seconds") @@ -100,11 +169,11 @@ fig = plt.figure(figsize=(8, 4)) ax1 = fig.add_subplot(121, projection='3d') ax2 = fig.add_subplot(122, projection='3d') -coils_optimized_initial_nearaxis.plot(ax=ax1, show=False) +init_coils.plot(ax=ax1, show=False) field_nearaxis_initial.plot(ax=ax1, show=False, alpha=0.35) tracing_initial.plot(ax=ax1, show=False) -coils_optimized.plot(ax=ax2, show=False) -field_nearaxis_optimized.plot(ax=ax2, show=False, alpha=0.35) +opt_coils.plot(ax=ax2, show=False) +opt_field_nearaxis.plot(ax=ax2, show=False, alpha=0.35) tracing_optimized.plot(ax=ax2, show=False) plt.show() @@ -115,7 +184,7 @@ # coils = Coils.from_json("stellarator_coils.json") # # Save results in vtk format to analyze in Paraview -# coils_optimized_initial_nearaxsis.to_vtk('coils_initial') -# coils_optimized.to_vtk('coils_optimized') +# init_coils.to_vtk('coils_initial') +# opt_coils.to_vtk('coils_optimized') # tracing_initial.to_vtk('trajectories_initial') # tracing_optimized.to_vtk('trajectories_final') \ No newline at end of file diff --git a/examples/coil_optimization/optimize_coils_and_surface.py b/examples/coil_optimization/optimize_coils_and_surface.py deleted file mode 100644 index f4a4f4e..0000000 --- a/examples/coil_optimization/optimize_coils_and_surface.py +++ /dev/null @@ -1,256 +0,0 @@ -import os -number_of_processors_to_use = 1 # Parallelization, this should divide ntheta*nphi -os.environ["XLA_FLAGS"] = f'--xla_force_host_platform_device_count={number_of_processors_to_use}' -from essos.fields import BiotSavart, near_axis -from essos.dynamics import Particles, Tracing -from essos.surfaces import BdotN_over_B, SurfaceRZFourier, B_on_surface -from essos.coils import Coils, CreateEquallySpacedCurves, Curves -from essos.optimization import optimize_loss_function, new_nearaxis_from_x_and_old_nearaxis -from essos.objective_functions import (loss_coil_curvature, difference_B_gradB_onaxis, - loss_coil_length,field_from_dofs) -import jax.numpy as jnp -from functools import partial -from jax import jit, vmap, devices, device_put, grad, debug -from jax.sharding import Mesh, NamedSharding, PartitionSpec -from time import time -import matplotlib.pyplot as plt - -mesh = Mesh(devices(), ("dev",)) -sharding = NamedSharding(mesh, PartitionSpec("dev", None)) - -ntheta=30 -nphi=30 -mpol=2 -ntor=2 -input = os.path.join(os.path.dirname(__file__), '..', 'input_files','input.rotating_ellipse') -surface_initial = SurfaceRZFourier.from_input_file(input, ntheta=ntheta, nphi=nphi, range_torus='half period') - -# Optimization parameters -max_coil_length = 38 -max_coil_curvature = 0.3 -order_Fourier_series_coils = 5 -number_coil_points = order_Fourier_series_coils*10 -maximum_function_evaluations = 20#600 -number_coils_per_half_field_period = 4 -tolerance_optimization = 1e-7 -target_B_on_axis = 5.7 - -nparticles = number_of_processors_to_use -maxtime_tracing = 4e-5 -num_steps=300 -trace_tolerance=1e-5 -model = 'GuidingCenterAdaptative' - -# Initialize coils -current_on_each_coil = 1.714e7 -number_of_field_periods = surface_initial.nfp -major_radius_coils = surface_initial.dofs[0] -minor_radius_coils = major_radius_coils/1.3 -curves = CreateEquallySpacedCurves(n_curves=number_coils_per_half_field_period, - order=order_Fourier_series_coils, - R=major_radius_coils, r=minor_radius_coils, - n_segments=number_coil_points, - nfp=number_of_field_periods, stellsym=True) -coils_initial = Coils(curves=curves, currents=[current_on_each_coil]*number_coils_per_half_field_period) - -# Initialize near-axis -rc=jnp.array([1, 0.045])*major_radius_coils -zs=jnp.array([0,-0.045])*major_radius_coils -etabar=-0.9/major_radius_coils -field_nearaxis_initial = near_axis(rc=rc, zs=zs, etabar=etabar, nfp=number_of_field_periods, B0=target_B_on_axis) - -print(f"Mean Magnetic field on surface: {jnp.mean(jnp.linalg.norm(B_on_surface(surface_initial, BiotSavart(coils_initial)), axis=2))}") - -# Initialize particles -# Xaxis = field_nearaxis_initial.R0*jnp.cos(field_nearaxis_initial.phi) -# Yaxis = field_nearaxis_initial.R0*jnp.sin(field_nearaxis_initial.phi) -# initial_xyz = jnp.array([Xaxis, Yaxis, field_nearaxis_initial.Z0]).T[:nparticles] -# particles = Particles(initial_xyz=initial_xyz, field=BiotSavart(coils_initial)) -# tracing_initial = Tracing(field=coils_initial, particles=particles, maxtime=maxtime_tracing, model=model, timesteps=num_steps) - -# # Plot initial state -# fig = plt.figure(figsize=(9, 8)) -# ax = fig.add_subplot(111, projection='3d') -# tracing_initial.plot(ax=ax, show=False) -# field_nearaxis_initial.plot(r=major_radius_coils/12, ax=ax, show=False) -# coils_initial.plot(ax=ax, show=False) -# surface_initial.plot(ax=ax, show=False) -# plt.show() - -# @partial(jit, static_argnames=['surface','field']) -def grad_AbsB_on_surface(surface, field): - ntheta = surface.ntheta - nphi = surface.nphi - gamma = surface.gamma - gamma_reshaped = gamma.reshape(nphi * ntheta, 3) - gamma_sharded = device_put(gamma_reshaped, sharding) - dAbsB_by_dX_on_surface = jit(vmap(field.dAbsB_by_dX), in_shardings=sharding, out_shardings=sharding)(gamma_sharded) - dAbsB_by_dX_on_surface = dAbsB_by_dX_on_surface.reshape(nphi, ntheta, 3) - return dAbsB_by_dX_on_surface - -# @partial(jit, static_argnames=['field']) -def B_dot_GradAbsB(points, field): - B = field.B(points) - GradAbsB = field.dAbsB_by_dX(points) - B_dot_GradAbsB = jnp.sum(B * GradAbsB, axis=-1) - return B_dot_GradAbsB - -# @partial(jit, static_argnames=['field']) -def grad_B_dot_GradAbsB(points, field): - return grad(B_dot_GradAbsB, argnums=0)(points, field) - -# @partial(jit, static_argnames=['surface','field']) -def grad_B_dot_GradAbsB_on_surface(surface, field): - ntheta = surface.ntheta - nphi = surface.nphi - gamma = surface.gamma - gamma_reshaped = gamma.reshape(nphi * ntheta, 3) - gamma_sharded = device_put(gamma_reshaped, sharding) - partial_grad_B_dot_GradAbsB = partial(grad_B_dot_GradAbsB, field=field) - grad_B_dot_GradAbsB_on_surface = jit(vmap(partial_grad_B_dot_GradAbsB), in_shardings=sharding, out_shardings=sharding)(gamma_sharded) - grad_B_dot_GradAbsB_on_surface = grad_B_dot_GradAbsB_on_surface.reshape(nphi, ntheta, 3) - return grad_B_dot_GradAbsB_on_surface - -# @partial(jit, static_argnames=['surface','field']) -def loss_normal_cross_GradB_dot_grad_B_dot_GradB_surface(surface, field): - gradAbsB_surface = grad_AbsB_on_surface(surface, field) - grad_B_dot_GradB_surface = grad_B_dot_GradAbsB_on_surface(surface, field) - normal_cross_GradB_surface = jnp.cross(surface.normal, gradAbsB_surface, axisa=-1, axisb=-1) - normal_cross_GradB_dot_grad_B_dot_GradB_surface = jnp.sum(normal_cross_GradB_surface * grad_B_dot_GradB_surface, axis=-1) - return normal_cross_GradB_dot_grad_B_dot_GradB_surface - -@partial(jit, static_argnums=(1, 5, 6, 7, 8, 9, 10)) -def loss_coils_and_surface(x, surface_all, field_nearaxis, dofs_curves, currents_scale, nfp, max_coil_length=42, - n_segments=60, stellsym=True, max_coil_curvature=0.5, target_B_on_surface=5.7): - - field=field_from_dofs(x[:-len(surface_all.x)-len(field_nearaxis.x)] ,dofs_curves=dofs_curves, currents_scale=currents_scale, nfp=nfp,n_segments=n_segments, stellsym=stellsym) - surface = SurfaceRZFourier(rc=surface_all.rc, zs=surface_all.zs, nfp=nfp, range_torus=surface_all.range_torus, nphi=surface_all.nphi, ntheta=surface_all.ntheta,mpol=surface_all.mpol,ntor=surface_all.ntor) - surface.dofs = x[-len(surface_all.x)-len(field_nearaxis.x):-len(field_nearaxis.x)] - - field_nearaxis = new_nearaxis_from_x_and_old_nearaxis(x[-len(field_nearaxis.x):], field_nearaxis) - - coil_length = field.coils.length - coil_curvature = field.coils.curvature - - - coil_length_loss = 1e3*jnp.max(jnp.maximum(0, coil_length - max_coil_length)) - coil_curvature_loss = 1e3*jnp.max(jnp.maximum(0, coil_curvature - max_coil_curvature)) - - normal_cross_GradB_dot_grad_B_dot_GradB_surface = jnp.sum(jnp.abs(loss_normal_cross_GradB_dot_grad_B_dot_GradB_surface(surface, field))) - - bdotn_over_b = BdotN_over_B(surface, field) - bdotn_over_b_loss = 10*jnp.sum(jnp.abs(bdotn_over_b)) - - mean_cross_sectional_area_loss = 100*jnp.abs(surface.mean_cross_sectional_area()-surface_all.mean_cross_sectional_area()) - - AbsB_on_surface = jnp.linalg.norm(B_on_surface(surface, field), axis=2) - AbsB_surface_loss = jnp.abs(jnp.mean(AbsB_on_surface)-target_B_on_surface) - - B_difference, gradB_difference = difference_B_gradB_onaxis(field_nearaxis, field) - B_difference_loss = 30*jnp.sum(jnp.abs(B_difference)) - gradB_difference_loss = 30*jnp.sum(jnp.abs(gradB_difference)) - - elongation = field_nearaxis.elongation - iota = field_nearaxis.iota - elongation_loss = jnp.sum(jnp.abs(elongation)) - iota_loss = 50/jnp.abs(iota) - - axis_surface = surface.dofs[0] - axis_nearaxis = field_nearaxis.rc[0] - axis_loss = jnp.abs(jnp.min(jnp.array([axis_nearaxis-axis_surface,0]))) - - # debug.print("######################") - # debug.print("normal_cross_GradB_dot_grad_B_dot_GradB_surface={}", normal_cross_GradB_dot_grad_B_dot_GradB_surface) - # debug.print("bdotn_over_b_loss={}", bdotn_over_b_loss) - # debug.print("mean_cross_sectional_area_loss={}", mean_cross_sectional_area_loss) - # debug.print("B_difference_loss={}", B_difference_loss) - # debug.print("gradB_difference_loss={}", gradB_difference_loss) - # debug.print("iota_loss={}", iota_loss) - # debug.print("axis_loss={}", axis_loss) - - # Xaxis = field_nearaxis.R0*jnp.cos(field_nearaxis.phi) - # Yaxis = field_nearaxis.R0*jnp.sin(field_nearaxis.phi) - # initial_xyz = jnp.array([Xaxis, Yaxis, field_nearaxis.Z0]).T[:nparticles] - # particles = Particles(initial_xyz=initial_xyz, field=field) - # particles_drift_loss = jnp.sum(loss_particle_drift(field, particles, maxtime_tracing, num_steps, trace_tolerance, model=model))/num_steps/nparticles - - return ( - coil_length_loss+coil_curvature_loss - # +normal_cross_GradB_dot_grad_B_dot_GradB_surface - +bdotn_over_b_loss - +mean_cross_sectional_area_loss - # +AbsB_surface_loss - +B_difference_loss - +gradB_difference_loss - +elongation_loss - +iota_loss - # +axis_loss - # +particles_drift_loss - ) - -# Optimize coils -print(f'Optimizing coils with {maximum_function_evaluations} function evaluations.') -time0 = time() -initial_dofs = jnp.concatenate((coils_initial.x, surface_initial.x, field_nearaxis_initial.x)) -coils_optimized, surface_optimized, field_nearaxis_optimized = optimize_loss_function(loss_coils_and_surface, initial_dofs=initial_dofs, coils=coils_initial, tolerance_optimization=tolerance_optimization, - maximum_function_evaluations=maximum_function_evaluations, surface_all=surface_initial, field_nearaxis=field_nearaxis_initial, - max_coil_length=max_coil_length, max_coil_curvature=max_coil_curvature, target_B_on_surface=target_B_on_axis) -print(f"Optimization took {time()-time0:.2f} seconds") -# Xaxis = field_nearaxis_optimized.R0*jnp.cos(field_nearaxis_optimized.phi) -# Yaxis = field_nearaxis_optimized.R0*jnp.sin(field_nearaxis_optimized.phi) -# initial_xyz = jnp.array([Xaxis, Yaxis, field_nearaxis_optimized.Z0]).T[:nparticles] -# particles = Particles(initial_xyz=initial_xyz, field=BiotSavart(coils_optimized)) -# tracing_optimized = Tracing(field=coils_optimized, particles=particles, maxtime=maxtime_tracing, model=model, timesteps=num_steps) - -print(f'############################################') -print(f"Mean Magnetic field on surface: {jnp.mean(jnp.linalg.norm(B_on_surface(surface_optimized, BiotSavart(coils_optimized)), axis=2))}") -print(f"Initial max(BdotN/B): {jnp.max(BdotN_over_B(surface_initial, BiotSavart(coils_initial))):.2e}") -print(f"Optimized max(BdotN/B): {jnp.max(BdotN_over_B(surface_optimized, BiotSavart(coils_optimized))):.2e}") -print(f'Initial iota on-axis: {field_nearaxis_initial.iota}') -print(f'Optimized iota on-axis: {field_nearaxis_optimized.iota}') -print(f'Initial max(elongation): {max(field_nearaxis_initial.elongation)}') -print(f'Optimized max(elongation): {max(field_nearaxis_optimized.elongation)}') -print(f"Initial coils length: {coils_initial.length[:number_coils_per_half_field_period]}") -print(f"Optimized coils length: {coils_optimized.length[:number_coils_per_half_field_period]}") -print(f"Initial coils curvature: {jnp.mean(coils_initial.curvature, axis=1)[:number_coils_per_half_field_period]}") -print(f"Optimized coils curvature: {jnp.mean(coils_optimized.curvature, axis=1)[:number_coils_per_half_field_period]}") -B_difference_initial, gradB_difference_initial = difference_B_gradB_onaxis(field_nearaxis_initial, BiotSavart(coils_initial)) -B_difference_optimized, gradB_difference_optimized = difference_B_gradB_onaxis(field_nearaxis_optimized, BiotSavart(coils_optimized)) -print(f'Initial B on axis difference: {jnp.sum(jnp.abs(B_difference_initial))}') -print(f'Optimized B on axis difference: {jnp.sum(jnp.abs(B_difference_optimized))}') -print(f'Initial gradB on axis difference: {jnp.sum(jnp.abs(gradB_difference_initial))}') -print(f'Optimized gradB on axis difference: {jnp.sum(jnp.abs(gradB_difference_optimized))}') - -# Plot coils, before and after optimization -fig = plt.figure(figsize=(8, 4)) -ax1 = fig.add_subplot(121, projection='3d') -ax2 = fig.add_subplot(122, projection='3d') -coils_initial.plot(ax=ax1, show=False) -surface_initial.plot(ax=ax1, show=False) -field_nearaxis_initial.plot(r=major_radius_coils/12, ax=ax1, show=False) -# tracing_initial.plot(ax=ax1, show=False) -coils_optimized.plot(ax=ax2, show=False) -surface_optimized.plot(ax=ax2, show=False) -field_nearaxis_optimized.plot(r=major_radius_coils/12, ax=ax2, show=False) -# tracing_optimized.plot(ax=ax2, show=False) -plt.tight_layout() -plt.show() -# Save the surface to a VMEC file -surface_optimized.to_vmec('input.optimized') - -# Save results in vtk format to analyze in Paraview -surface_initial.to_vtk('initial_surface', field=BiotSavart(coils_initial)) -coils_initial.to_vtk('initial_coils') -field_nearaxis_initial.to_vtk('initial_field_nearaxis', r=major_radius_coils/12, field=BiotSavart(coils_initial)) -surface_optimized.to_vtk('optimized_surface', field=BiotSavart(coils_optimized)) -coils_optimized.to_vtk('optimized_coils') -field_nearaxis_optimized.to_vtk('optimized_field_nearaxis', r=major_radius_coils/12, field=BiotSavart(coils_optimized)) - -# tracing_initial.to_vtk('initial_tracing') -# tracing_optimized.to_vtk('optimized_tracing') - -# # Save the coils to a json file -# coils_optimized.to_json("stellarator_coils.json") -# # Load the coils from a json file -# from essos.coils import Coils -# coils = Coils.from_json("stellarator_coils.json") \ No newline at end of file diff --git a/examples/coil_optimization/optimize_coils_for_nearaxis.py b/examples/coil_optimization/optimize_coils_for_nearaxis.py index a77ce7e..f32eaa5 100644 --- a/examples/coil_optimization/optimize_coils_for_nearaxis.py +++ b/examples/coil_optimization/optimize_coils_for_nearaxis.py @@ -1,3 +1,6 @@ +import os +number_of_processors_to_use = 4 # Parallelization, this should divide nfieldlines +os.environ["XLA_FLAGS"] = f'--xla_force_host_platform_device_count={number_of_processors_to_use}' from time import time import jax.numpy as jnp import matplotlib.pyplot as plt @@ -7,45 +10,118 @@ from essos.dynamics import Tracing from essos.optimization import optimize_loss_function -from essos.objective_functions import loss_coils_for_nearaxis - -# Optimization parameters -max_coil_length = 4 -max_coil_curvature = 6 -order_Fourier_series_coils = 5 -number_coil_points = order_Fourier_series_coils*10 -maximum_function_evaluations = 200 -number_coils_per_half_field_period = 3 +from jax import vmap, jit +# In this exmple, `scipy.optimize.least_squares` is used, but any other optimizer, e.g. from +# `scipy.optimize.minimize` or `jaxopt`, can be used as well and may even be preferable. +from scipy.optimize import least_squares +from essos.losses import custom_loss + + +""" Creating starting coils and surface """ +N_COILS = 3; FOURIER_ORDER = 6; LARGE_R = 10; SMALL_R = 5.6; NFP = 3; N_SEGMENTS = 60; STELLSYM = True # Curve parameters +COIL_CURRENT = 1. # Amperes (optimization does not depend on current magnitude) tolerance_optimization = 1e-8 +maximum_function_evaluations =200 + + # Initialize Near-Axis field -rc=jnp.array([1, 0.045]) -zs=jnp.array([0,-0.045]) +rc=jnp.array([1., 0.045]) +zs=jnp.array([0.,-0.045]) etabar=-0.9 -nfp=3 -field = near_axis(rc=rc, zs=zs, etabar=etabar, nfp=nfp) +field_nearaxis_initial = near_axis(rc=rc, zs=zs, etabar=etabar, nfp=NFP,order='r3') # Initialize coils -current_on_each_coil = 17e5*field.B0/nfp/2 -number_of_field_periods = nfp -major_radius_coils = field.R0[0] +current_on_each_coil = 17.e5*field_nearaxis_initial.B0/NFP/2. +number_of_field_periods = NFP +major_radius_coils = field_nearaxis_initial.R0[0] minor_radius_coils = major_radius_coils/2.0 -curves = CreateEquallySpacedCurves(n_curves=number_coils_per_half_field_period, - order=order_Fourier_series_coils, +init_curves = CreateEquallySpacedCurves(n_curves=N_COILS, + order=FOURIER_ORDER, R=major_radius_coils, r=minor_radius_coils, - n_segments=number_coil_points, - nfp=number_of_field_periods, stellsym=True) -coils_initial = Coils(curves=curves, currents=[current_on_each_coil]*number_coils_per_half_field_period) + n_segments=N_SEGMENTS, + nfp=number_of_field_periods, stellsym=STELLSYM) +init_coils = Coils(curves=init_curves, currents=jnp.array([current_on_each_coil]*N_COILS)) +init_field = BiotSavart(init_coils) + + +""" Setting the losses weights and targets """ +LENGTH_WEIGHT = 1.; LENGTH_TARGET = 4. +CURVATURE_WEIGHT = 1.; CURVATURE_TARGET = 6. +B_DIFFERENCE_WEIGHT = 1. +GRADB_DIFFERENCE_WEIGHT = 1. + + + +""" Creating the loss functions """ +def near_axis_field_quantities(field_nearaxis): + Raxis = field_nearaxis.R0 + Zaxis = field_nearaxis.Z0 + phi = field_nearaxis.phi + Xaxis = Raxis*jnp.cos(phi) + Yaxis = Raxis*jnp.sin(phi) + points = jnp.array([Xaxis, Yaxis, Zaxis]) + B_nearaxis = field_nearaxis.B_axis.T + gradB_nearaxis = field_nearaxis.grad_B_axis.T + return points, B_nearaxis, gradB_nearaxis + + +def loss_B_difference_coils_near_axis(field, field_nearaxis): + points, B_nearaxis, _ = near_axis_field_quantities(field_nearaxis) + B_coils = vmap(field.B)(points.T) + B_difference_loss = jnp.sum(jnp.abs(jnp.array(B_coils)-jnp.array(B_nearaxis))) + return B_difference_loss + +def loss_gradB_difference_coils_near_axis(field, field_nearaxis): + points, _, gradB_nearaxis = near_axis_field_quantities(field_nearaxis) + gradB_coils = vmap(field.dB_by_dX)(points.T) + gradB_difference_loss = jnp.sum(jnp.abs(jnp.array(gradB_coils)-jnp.array(gradB_nearaxis))) + return gradB_difference_loss + +def loss_length(field,length_target=LENGTH_TARGET): + return jnp.mean(jnp.maximum(0, field.coils.length - length_target)) + +def loss_curvature(field,curvature_target=CURVATURE_TARGET): + return jnp.mean(jnp.maximum(0, field.coils.curvature - curvature_target)) + + +""" Defining custom losses """ +L_B_difference = custom_loss(loss_B_difference_coils_near_axis, "field", field_nearaxis=field_nearaxis_initial) +L_gradB_difference = custom_loss(loss_gradB_difference_coils_near_axis, "field", field_nearaxis=field_nearaxis_initial) +L_length = custom_loss(loss_length, "field") +L_curvature = custom_loss(loss_curvature, "field") +""" Defining total loss + setting dependencies """ +L_total = B_DIFFERENCE_WEIGHT*L_B_difference + GRADB_DIFFERENCE_WEIGHT*L_gradB_difference + LENGTH_WEIGHT*L_length + CURVATURE_WEIGHT*L_curvature + + +L_total.dependencies = {"field": init_field, "field_nearaxis": field_nearaxis_initial} + +""" Optimizing the total loss """ +t_start = time() +res = least_squares(L_total, L_total.starting_dofs, L_total.grad, verbose=2, ftol=1e-5, gtol=1e-5, xtol=1e-14, max_nfev=200) +t_end = time() + +print(f"\nOptimization took {t_end - t_start:.2f} seconds") +print("Initial loss:", L_total(L_total.starting_dofs)) +print("Loss after optimization:", L_total(res.x)) + +opt_field = L_total.dofs_to_pytree(res.x)["field"] +opt_coils = opt_field.coils + + +B_difference_initial = loss_B_difference_coils_near_axis(init_field, field_nearaxis_initial) +gradB_difference_initial = loss_gradB_difference_coils_near_axis(init_field, field_nearaxis_initial) + +B_difference_optimized = loss_B_difference_coils_near_axis(opt_field, field_nearaxis_initial) +gradB_difference_optimized = loss_gradB_difference_coils_near_axis(opt_field, field_nearaxis_initial) + + +print(f'############################################') +print(f'Loss of B difference for initial near-axis: {B_difference_initial}') +print(f'Loss of B difference for optimized near-axis: {B_difference_optimized}') +print(f'Loss of gradB difference for initial near-axis: {gradB_difference_initial}') +print(f'Loss of gradB difference for optimized near-axis: {gradB_difference_optimized}') -# Optimize coils -print(f'Optimizing coils with {maximum_function_evaluations} function evaluations.') -time0 = time() -initial_dofs = coils_initial.x -coils_optimized = optimize_loss_function(loss_coils_for_nearaxis, initial_dofs=coils_initial.x, - coils=coils_initial, tolerance_optimization=tolerance_optimization, - maximum_function_evaluations=maximum_function_evaluations, field_nearaxis=field, - max_coil_length=max_coil_length, max_coil_curvature=max_coil_curvature,) -print(f"Optimization took {time()-time0:.2f} seconds") # Trace fieldlines nfieldlines = 6 @@ -53,15 +129,17 @@ tmax = 1.e-6 trace_tolerance = 1e-7 -R0 = jnp.linspace(field.R0[0], 1.05*field.R0[0], nfieldlines) +R0_initial = jnp.linspace(field_nearaxis_initial.R0[0], 1.05*field_nearaxis_initial.R0[0], nfieldlines) +R0_optimized = jnp.linspace(field_nearaxis_initial.R0[0], 1.05*field_nearaxis_initial.R0[0], nfieldlines) Z0 = jnp.zeros(nfieldlines) phi0 = jnp.zeros(nfieldlines) -initial_xyz=jnp.array([R0*jnp.cos(phi0), R0*jnp.sin(phi0), Z0]).T +initial_xyz_initial = jnp.array([R0_initial*jnp.cos(phi0), R0_initial*jnp.sin(phi0), Z0]).T +initial_xyz_optimized = jnp.array([R0_optimized*jnp.cos(phi0), R0_optimized*jnp.sin(phi0), Z0]).T time0 = time() -tracing_initial = Tracing(field=BiotSavart(coils_initial), model='FieldLineAdaptative', initial_conditions=initial_xyz, +tracing_initial = Tracing(field=init_field, model='FieldLineAdaptative', initial_conditions=initial_xyz_initial, maxtime=tmax, times_to_trace=num_steps, atol=trace_tolerance,rtol=trace_tolerance) -tracing_optimized = Tracing(field=BiotSavart(coils_optimized), model='FieldLineAdaptative', initial_conditions=initial_xyz, +tracing_optimized = Tracing(field=opt_field, model='FieldLineAdaptative', initial_conditions=initial_xyz_optimized, maxtime=tmax, times_to_trace=num_steps, atol=trace_tolerance,rtol=trace_tolerance) print(f"Tracing took {time()-time0:.2f} seconds") @@ -69,11 +147,11 @@ fig = plt.figure(figsize=(8, 4)) ax1 = fig.add_subplot(121, projection='3d') ax2 = fig.add_subplot(122, projection='3d') -coils_initial.plot(ax=ax1, show=False) -field.plot(ax=ax1, show=False, alpha=0.2) +init_coils.plot(ax=ax1, show=False) +field_nearaxis_initial.plot(ax=ax1, show=False, alpha=0.35) tracing_initial.plot(ax=ax1, show=False) -coils_optimized.plot(ax=ax2, show=False) -field.plot(ax=ax2, show=False, alpha=0.2) +opt_coils.plot(ax=ax2, show=False) +field_nearaxis_initial.plot(ax=ax2, show=False, alpha=0.35) tracing_optimized.plot(ax=ax2, show=False) plt.show() @@ -84,7 +162,7 @@ # coils = Coils.from_json("stellarator_coils.json") # # Save results in vtk format to analyze in Paraview -# coils_initial.to_vtk('coils_initial') -# coils_optimized.to_vtk('coils_optimized') +# init_coils.to_vtk('coils_initial') +# opt_coils.to_vtk('coils_optimized') # tracing_initial.to_vtk('trajectories_initial') # tracing_optimized.to_vtk('trajectories_final') \ No newline at end of file diff --git a/examples/coil_optimization/optimize_coils_particle_confinement_fullorbit.py b/examples/coil_optimization/optimize_coils_particle_confinement_fullorbit.py index 27ca0a6..d68d3a1 100644 --- a/examples/coil_optimization/optimize_coils_particle_confinement_fullorbit.py +++ b/examples/coil_optimization/optimize_coils_particle_confinement_fullorbit.py @@ -7,65 +7,119 @@ import matplotlib.pyplot as plt from essos.dynamics import Particles, Tracing from essos.coils import Coils, CreateEquallySpacedCurves +from jax import vmap, jit +# In this exmple, `scipy.optimize.least_squares` is used, but any other optimizer, e.g. from +# `scipy.optimize.minimize` or `jaxopt`, can be used as well and may even be preferable. +from scipy.optimize import least_squares +from essos.losses import custom_loss from essos.fields import BiotSavart -from essos.optimization import optimize_loss_function -from essos.objective_functions import loss_optimize_coils_for_particle_confinement -from essos.fields import BiotSavart + +# Particle optimization parameters # Optimization parameters -target_B_on_axis = 5.7 -max_coil_length = 31 -max_coil_curvature = 0.4 -nparticles = number_of_processors_to_use*1 -order_Fourier_series_coils = 4 -number_coil_points = 80 -maximum_function_evaluations = 10 -maxtime_tracing = 1e-6 -number_coils_per_half_field_period = 3 -number_of_field_periods = 2 -model = 'FullOrbit_Boris' -timesteps = 3000#int(3*maxtime_tracing/1e-8) - -nparticles_plot = number_of_processors_to_use*2 -model_plot = 'GuidingCenterAdaptative' -timesteps_plot = 10000 -maxtime_tracing_plot = 3e-5 +NPARTICLES = number_of_processors_to_use*10 +MAXTIME_TRACING = 1e-4 +NUMBER_COILS_PER_HALF_FIELD_PERIOD = 3 +NUMBER_OF_FIELD_PERIODS = 2 +MODEL = 'FullOrbit_Boris' +TIMESTEP=1.e-14 +TRACE_TOLERANCE=1e-8 +NUM_STEPS=1000 + + +NPARTICLES_PLOT = number_of_processors_to_use*10 +MAXTIME_TRACING_PLOT = 1e-4 + +""" Creating starting coils and surface """ +N_COILS = 3 +FOURIER_ORDER = 6 +LARGE_R = 7.74 +SMALL_R = 4.5 +NFP = 2 +N_SEGMENTS = 60 +STELLSYM = True # Curve parameters +COIL_CURRENT = 1.84e7 # Amperes (optimization does not depend on current magnitude) +MAXIMUM_FUNCTION_EVALUATIONS =200 # Initialize coils -current_on_each_coil = 1.84e7 -major_radius_coils = 7.75 -minor_radius_coils = 4.5 -curves = CreateEquallySpacedCurves(n_curves=number_coils_per_half_field_period, - order=order_Fourier_series_coils, - R=major_radius_coils, r=minor_radius_coils, - n_segments=number_coil_points, - nfp=number_of_field_periods, stellsym=True) -coils_initial = Coils(curves=curves, currents=[current_on_each_coil]*number_coils_per_half_field_period) +init_curves = CreateEquallySpacedCurves(n_curves=N_COILS, + order=FOURIER_ORDER, + R=LARGE_R, r=SMALL_R, + n_segments=N_SEGMENTS, + nfp=NFP, stellsym=STELLSYM) +init_coils = Coils(curves=init_curves, currents=jnp.array([COIL_CURRENT]*N_COILS)) +init_field = BiotSavart(init_coils) + +""" Setting the losses weights and targets """ +LENGTH_WEIGHT = 1.; LENGTH_TARGET = 31. +CURVATURE_WEIGHT = 1.; CURVATURE_TARGET = 0.4 +BAXIS_WEIGHT = 1.; BAXIS_TARGET = 5.7 +RADIAL_DRIFT_WEIGHT = 1. + +def loss_particle_radial_drift(field, particles, timestep=1.e-8, maxtime=1e-5, num_steps=300, trace_tolerance=1e-5, model='GuidingCenterAdaptative',boundary=None): + particles.to_full_orbit(field) + 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) + xyz = tracing.trajectories[:,:, :3] + R_axis=field.r_axis + Z_axis=field.z_axis + #Ideally here one would differentiate in time through diffrax !TODO + r_cross=jnp.sqrt(jnp.square(jnp.sqrt(jnp.square(xyz[:,0])+jnp.square(xyz[:,1]))-R_axis+1.e-12)+jnp.square(xyz[:,2]-Z_axis+1.e-12)) + v_r_cross=jnp.diff(r_cross,axis=1)#/tracing.times_to_trace*tracing.maxtime + return (jnp.sum(jnp.square(jnp.average(v_r_cross,axis=1)))) + +def normB_axis(field, npoints=15): + R_axis=field.r_axis + phi_array = jnp.linspace(0, 2 * jnp.pi, npoints) + B_axis = vmap(lambda phi: field.AbsB(jnp.array([R_axis * jnp.cos(phi), R_axis * jnp.sin(phi), 0])))(phi_array) + return B_axis + +def loss_normB_axis_average(field,npoints=15, target_B=BAXIS_TARGET): + B_axis = normB_axis(field, npoints) + return jnp.abs(jnp.average(B_axis)-target_B) + +def loss_length(field,length_target=LENGTH_TARGET): + return jnp.mean(jnp.maximum(0, field.coils.length - length_target)) + +def loss_curvature(field,curvature_target=CURVATURE_TARGET): + return jnp.mean(jnp.maximum(0, field.coils.curvature - curvature_target)) # Initialize particles -phi_array = jnp.linspace(0, 2*jnp.pi, nparticles) -initial_xyz=jnp.array([major_radius_coils*jnp.cos(phi_array), major_radius_coils*jnp.sin(phi_array), 0*phi_array]).T +phi_array = jnp.linspace(0, 2*jnp.pi, NPARTICLES) +initial_xyz=jnp.array([LARGE_R*jnp.cos(phi_array), LARGE_R*jnp.sin(phi_array), 0*phi_array]).T particles = Particles(initial_xyz=initial_xyz) -particles.to_full_orbit(BiotSavart(coils_initial)) -tracing_initial = Tracing(field=coils_initial, particles=particles, maxtime=maxtime_tracing, model=model, times_to_trace=timesteps) - -# Optimize coils -print(f'Optimizing coils with {maximum_function_evaluations} function evaluations and maxtime_tracing={maxtime_tracing}') -time0 = time() -coils_optimized = optimize_loss_function(loss_optimize_coils_for_particle_confinement, initial_dofs=coils_initial.x, coils=coils_initial, - tolerance_optimization=1e-4, particles=particles, maximum_function_evaluations=maximum_function_evaluations, - max_coil_curvature=max_coil_curvature, target_B_on_axis=target_B_on_axis, max_coil_length=max_coil_length, - model=model, maxtime=maxtime_tracing, num_steps=500, trace_tolerance=1e-5) -# coils_optimized = optimize_coils_for_particle_confinement(coils_initial, particles, target_B_on_axis=target_B_on_axis, maxtime=maxtime_tracing, model=model, -# max_coil_length=max_coil_length, maximum_function_evaluations=maximum_function_evaluations, max_coil_curvature=max_coil_curvature) -print(f" Optimization took {time()-time0:.2f} seconds") -particles.to_full_orbit(BiotSavart(coils_optimized)) - -phi_array_plot = jnp.linspace(0, 2*jnp.pi, nparticles_plot) -initial_xyz_plot=jnp.array([major_radius_coils*jnp.cos(phi_array_plot), major_radius_coils*jnp.sin(phi_array_plot), 0*phi_array_plot]).T + + +""" Defining custom losses """ +L_radial_drift= custom_loss(loss_particle_radial_drift, "field", particles=particles, timestep=TIMESTEP, maxtime=MAXTIME_TRACING, num_steps=NUM_STEPS, trace_tolerance=TRACE_TOLERANCE, model=MODEL) +L_B_axis= custom_loss(loss_normB_axis_average, "field") +L_length = custom_loss(loss_length, "field") +L_curvature = custom_loss(loss_curvature, "field") +""" Defining total loss + setting dependencies """ +L_total = RADIAL_DRIFT_WEIGHT*L_radial_drift + L_B_axis + LENGTH_WEIGHT*L_length + CURVATURE_WEIGHT*L_curvature + + +L_total.dependencies = {"field": init_field} + +""" Optimizing the total loss """ +t_start = time() +res = least_squares(L_total, L_total.starting_dofs, L_total.grad, verbose=2, ftol=1e-5, gtol=1e-5, xtol=1e-14, max_nfev=MAXIMUM_FUNCTION_EVALUATIONS) +t_end = time() + +print(f"\nOptimization took {t_end - t_start:.2f} seconds") +print("Initial loss:", L_total(L_total.starting_dofs)) +print("Loss after optimization:", L_total(res.x)) + +opt_field = L_total.dofs_to_pytree(res.x)["field"] +opt_coils = opt_field.coils + +""" Plotting results """ +phi_array_plot = jnp.linspace(0, 2*jnp.pi, NPARTICLES_PLOT) +initial_xyz_plot=jnp.array([LARGE_R*jnp.cos(phi_array_plot), LARGE_R*jnp.sin(phi_array_plot), 0*phi_array_plot]).T particles_plot = Particles(initial_xyz=initial_xyz_plot) -particles.to_full_orbit(BiotSavart(coils_optimized)) -tracing_optimized = Tracing(field=coils_optimized, particles=particles, maxtime=maxtime_tracing_plot, model=model_plot, times_to_trace=timesteps_plot) +particles_plot.to_full_orbit(opt_field) +tracing_initial = Tracing(field=init_field, particles=particles_plot, maxtime=MAXTIME_TRACING_PLOT, model=MODEL, times_to_trace=NUM_STEPS, timestep=TIMESTEP, atol=TRACE_TOLERANCE, rtol=TRACE_TOLERANCE) +tracing_optimized = Tracing(field=opt_field, particles=particles_plot, maxtime=MAXTIME_TRACING_PLOT, model=MODEL, times_to_trace=NUM_STEPS, timestep=TIMESTEP, atol=TRACE_TOLERANCE, rtol=TRACE_TOLERANCE) # Plot trajectories, before and after optimization fig = plt.figure(figsize=(9, 8)) @@ -74,15 +128,17 @@ ax3 = fig.add_subplot(223) ax4 = fig.add_subplot(224) -coils_initial.plot(ax=ax1, show=False) +init_coils.plot(ax=ax1, show=False) tracing_initial.plot(ax=ax1, show=False) for i, trajectory in enumerate(tracing_initial.trajectories): ax3.plot(jnp.sqrt(trajectory[:,0]**2+trajectory[:,1]**2), trajectory[:, 2], label=f'Particle {i+1}') + ax3.set_xlabel('R (m)');ax3.set_ylabel('Z (m)');#ax3.legend() -coils_optimized.plot(ax=ax2, show=False) +opt_coils.plot(ax=ax2, show=False) tracing_optimized.plot(ax=ax2, show=False) for i, trajectory in enumerate(tracing_optimized.trajectories): ax4.plot(jnp.sqrt(trajectory[:,0]**2+trajectory[:,1]**2), trajectory[:, 2], label=f'Particle {i+1}') + ax4.set_xlabel('R (m)');ax4.set_ylabel('Z (m)');#ax4.legend() plt.tight_layout() plt.show() diff --git a/examples/coil_optimization/optimize_coils_particle_confinement_guidingcenter.py b/examples/coil_optimization/optimize_coils_particle_confinement_guidingcenter.py new file mode 100644 index 0000000..646f556 --- /dev/null +++ b/examples/coil_optimization/optimize_coils_particle_confinement_guidingcenter.py @@ -0,0 +1,157 @@ + +import os +number_of_processors_to_use = 1 # Parallelization, this should divide nparticles +os.environ["XLA_FLAGS"] = f'--xla_force_host_platform_device_count={number_of_processors_to_use}' +from time import time +import jax.numpy as jnp +import matplotlib.pyplot as plt +from essos.dynamics import Particles, Tracing +from essos.coils import Coils, CreateEquallySpacedCurves +from jax import vmap, jit +# In this exmple, `scipy.optimize.least_squares` is used, but any other optimizer, e.g. from +# `scipy.optimize.minimize` or `jaxopt`, can be used as well and may even be preferable. +from scipy.optimize import least_squares +from essos.losses import custom_loss +from essos.fields import BiotSavart + +# Particle optimization parameters +# Optimization parameters +NPARTICLES = number_of_processors_to_use*10 +MAXTIME_TRACING = 1e-4 +NUMBER_COILS_PER_HALF_FIELD_PERIOD = 3 +NUMBER_OF_FIELD_PERIODS = 2 +MODEL = 'GuidingCenterAdaptative' +TIMESTEP=1.e-14 +TRACE_TOLERANCE=1e-8 +NUM_STEPS=1000 + +NPARTICLES_PLOT = number_of_processors_to_use*10 +MAXTIME_TRACING_PLOT = 1e-4 + + + +""" Creating starting coils and surface """ +N_COILS = 3 +FOURIER_ORDER = 6 +LARGE_R = 7.74 +SMALL_R = 4.5 +NFP = 2 +N_SEGMENTS = 60 +STELLSYM = True # Curve parameters +COIL_CURRENT = 1.84e7 # Amperes (optimization does not depend on current magnitude) +MAXIMUM_FUNCTION_EVALUATIONS =200 + +# Initialize coils +init_curves = CreateEquallySpacedCurves(n_curves=N_COILS, + order=FOURIER_ORDER, + R=LARGE_R, r=SMALL_R, + n_segments=N_SEGMENTS, + nfp=NFP, stellsym=STELLSYM) +init_coils = Coils(curves=init_curves, currents=jnp.array([COIL_CURRENT]*N_COILS)) +init_field = BiotSavart(init_coils) + +""" Setting the losses weights and targets """ +LENGTH_WEIGHT = 1.; LENGTH_TARGET = 31. +CURVATURE_WEIGHT = 1.; CURVATURE_TARGET = 0.4 +BAXIS_WEIGHT = 1.; BAXIS_TARGET = 5.7 +RADIAL_DRIFT_WEIGHT = 1. + +def loss_particle_radial_drift(field, particles, timestep=1.e-8, maxtime=1e-5, num_steps=300, trace_tolerance=1e-5, model='GuidingCenterAdaptative',boundary=None): + particles.to_full_orbit(field) + 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) + xyz = tracing.trajectories[:,:, :3] + R_axis=field.r_axis + Z_axis=field.z_axis + #Ideally here one would differentiate in time through diffrax !TODO + r_cross=jnp.sqrt(jnp.square(jnp.sqrt(jnp.square(xyz[:,:,0])+jnp.square(xyz[:,:,1]))-R_axis+1.e-12)+jnp.square(xyz[:,:,2]-Z_axis+1.e-12)) + v_r_cross=jnp.diff(r_cross,axis=1)#/tracing.times_to_trace*tracing.maxtime + return (jnp.sum(jnp.square(jnp.average(v_r_cross,axis=1)))) + +def normB_axis(field, npoints=15): + R_axis=field.r_axis + phi_array = jnp.linspace(0, 2 * jnp.pi, npoints) + B_axis = vmap(lambda phi: field.AbsB(jnp.array([R_axis * jnp.cos(phi), R_axis * jnp.sin(phi), 0])))(phi_array) + return B_axis + +def loss_normB_axis_average(field,npoints=15, target_B=BAXIS_TARGET): + B_axis = normB_axis(field, npoints) + return jnp.abs(jnp.average(B_axis)-target_B) + +def loss_length(field,length_target=LENGTH_TARGET): + return jnp.mean(jnp.maximum(0, field.coils.length - length_target)) + +def loss_curvature(field,curvature_target=CURVATURE_TARGET): + return jnp.mean(jnp.maximum(0, field.coils.curvature - curvature_target)) + + + +# Initialize particles +phi_array = jnp.linspace(0, 2*jnp.pi, NPARTICLES) +initial_xyz=jnp.array([LARGE_R*jnp.cos(phi_array), LARGE_R*jnp.sin(phi_array), 0*phi_array]).T +particles = Particles(initial_xyz=initial_xyz) + + +""" Defining custom losses """ +L_radial_drift= custom_loss(loss_particle_radial_drift, "field", particles=particles, timestep=TIMESTEP, maxtime=MAXTIME_TRACING, num_steps=NUM_STEPS, trace_tolerance=TRACE_TOLERANCE, model=MODEL) +L_B_axis= custom_loss(loss_normB_axis_average, "field") +L_length = custom_loss(loss_length, "field") +L_curvature = custom_loss(loss_curvature, "field") +""" Defining total loss + setting dependencies """ +L_total = RADIAL_DRIFT_WEIGHT*L_radial_drift + L_B_axis + LENGTH_WEIGHT*L_length + CURVATURE_WEIGHT*L_curvature + +L_total.dependencies = {"field": init_field} + +""" Optimizing the total loss """ +t_start = time() +res = least_squares(L_total, L_total.starting_dofs, L_total.grad, verbose=2, ftol=1e-5, gtol=1e-5, xtol=1e-14, max_nfev=MAXIMUM_FUNCTION_EVALUATIONS) +t_end = time() + +print(f"\nOptimization took {t_end - t_start:.2f} seconds") +print("Initial loss:", L_total(L_total.starting_dofs)) +print("Loss after optimization:", L_total(res.x)) + +opt_field = L_total.dofs_to_pytree(res.x)["field"] +opt_coils = opt_field.coils + +""" Plotting results """ +phi_array_plot = jnp.linspace(0, 2*jnp.pi, NPARTICLES_PLOT) +initial_xyz_plot=jnp.array([LARGE_R*jnp.cos(phi_array_plot), LARGE_R*jnp.sin(phi_array_plot), 0*phi_array_plot]).T +particles_plot = Particles(initial_xyz=initial_xyz_plot) +particles_plot.to_full_orbit(opt_field) +tracing_initial = Tracing(field=init_field, particles=particles_plot, maxtime=MAXTIME_TRACING_PLOT, model=MODEL, times_to_trace=NUM_STEPS, timestep=TIMESTEP, atol=TRACE_TOLERANCE, rtol=TRACE_TOLERANCE) +tracing_optimized = Tracing(field=opt_field, particles=particles_plot, maxtime=MAXTIME_TRACING_PLOT, model=MODEL, times_to_trace=NUM_STEPS, timestep=TIMESTEP, atol=TRACE_TOLERANCE, rtol=TRACE_TOLERANCE) + +# Plot trajectories, before and after optimization +fig = plt.figure(figsize=(9, 8)) +ax1 = fig.add_subplot(221, projection='3d') +ax2 = fig.add_subplot(222, projection='3d') +ax3 = fig.add_subplot(223) +ax4 = fig.add_subplot(224) + +init_coils.plot(ax=ax1, show=False) +tracing_initial.plot(ax=ax1, show=False) +for i, trajectory in enumerate(tracing_initial.trajectories): + ax3.plot(jnp.sqrt(trajectory[:,0]**2+trajectory[:,1]**2), trajectory[:, 2], label=f'Particle {i+1}') + +ax3.set_xlabel('R (m)');ax3.set_ylabel('Z (m)');#ax3.legend() +opt_coils.plot(ax=ax2, show=False) +tracing_optimized.plot(ax=ax2, show=False) +for i, trajectory in enumerate(tracing_optimized.trajectories): + ax4.plot(jnp.sqrt(trajectory[:,0]**2+trajectory[:,1]**2), trajectory[:, 2], label=f'Particle {i+1}') + +ax4.set_xlabel('R (m)');ax4.set_ylabel('Z (m)');#ax4.legend() +plt.tight_layout() +plt.show() + +# # Save the coils to a json file +# coils_optimized.to_json("stellarator_coils.json") +# # Load the coils from a json file +# from essos.coils import Coils_from_json +# coils = Coils_from_json("stellarator_coils.json") + +# # Save results in vtk format to analyze in Paraview +# tracing_initial.to_vtk('trajectories_initial') +# tracing_optimized.to_vtk('trajectories_final') +# coils_initial.to_vtk('coils_initial') +# coils_optimized.to_vtk('coils_optimized') diff --git a/examples/coil_optimization/optimize_coils_particle_confinement_guidingcenter_adam.py b/examples/coil_optimization/optimize_coils_particle_confinement_guidingcenter_adam.py deleted file mode 100644 index 56ee3e9..0000000 --- a/examples/coil_optimization/optimize_coils_particle_confinement_guidingcenter_adam.py +++ /dev/null @@ -1,125 +0,0 @@ - -import os -number_of_processors_to_use = 1 # Parallelization, this should divide nparticles -os.environ["XLA_FLAGS"] = f'--xla_force_host_platform_device_count={number_of_processors_to_use}' -import jax.numpy as jnp -from jax import jit, grad -import matplotlib.pyplot as plt -from essos.dynamics import Particles, Tracing -from essos.coils import Coils, CreateEquallySpacedCurves,Curves -from essos.objective_functions import loss_particle_r_cross_max_constraint -from essos.objective_functions import loss_coil_curvature_new as loss_coil_curvature, loss_coil_length_new as loss_coil_length, loss_normB_axis_average -from functools import partial -import optax - - -# Optimization parameters -target_B_on_axis = 5.7 -max_coil_length = 31 -max_coil_curvature = 0.4 -nparticles = number_of_processors_to_use*1 -order_Fourier_series_coils = 4 -number_coil_points = 80 -maximum_function_evaluations = 15 -maxtimes = [2.e-5] -num_steps=100 -number_coils_per_half_field_period = 3 -number_of_field_periods = 2 -model = 'GuidingCenterAdaptative' - -# Initialize coils -current_on_each_coil = 1.84e7 -major_radius_coils = 7.75 -minor_radius_coils = 4.45 -curves = CreateEquallySpacedCurves(n_curves=number_coils_per_half_field_period, - order=order_Fourier_series_coils, - R=major_radius_coils, r=minor_radius_coils, - n_segments=number_coil_points, - nfp=number_of_field_periods, stellsym=True) -coils_initial = Coils(curves=curves, currents=[current_on_each_coil]*number_coils_per_half_field_period) - -len_dofs_curves = len(jnp.ravel(coils_initial.dofs_curves)) -nfp = coils_initial.nfp -stellsym = coils_initial.stellsym -n_segments = coils_initial.n_segments -dofs_curves_shape = coils_initial.dofs_curves.shape -currents_scale = coils_initial.currents_scale - -# Initialize particles -phi_array = jnp.linspace(0, 2*jnp.pi, nparticles) -initial_xyz=jnp.array([major_radius_coils*jnp.cos(phi_array), major_radius_coils*jnp.sin(phi_array), 0*phi_array]).T -particles = Particles(initial_xyz=initial_xyz) - -t=maxtimes[0] - -curvature_partial=partial(loss_coil_curvature, dofs_curves=coils_initial.dofs_curves, currents_scale=currents_scale, nfp=nfp, n_segments=n_segments, stellsym=stellsym,max_coil_curvature=max_coil_curvature) -length_partial=partial(loss_coil_length, dofs_curves=coils_initial.dofs_curves, currents_scale=currents_scale, nfp=nfp, n_segments=n_segments, stellsym=stellsym,max_coil_length=max_coil_length) -Baxis_average_partial=partial(loss_normB_axis_average,dofs_curves=coils_initial.dofs_curves, currents_scale=currents_scale, nfp=nfp, n_segments=n_segments, stellsym=stellsym,npoints=15,target_B_on_axis=target_B_on_axis) -r_max_partial = partial(loss_particle_r_cross_max_constraint, particles=particles,dofs_curves=coils_initial.dofs_curves, currents_scale=currents_scale, nfp=nfp, n_segments=n_segments, stellsym=stellsym,maxtime=t,model = model,num_steps=num_steps) - -params=coils_initial.x -optimizer=optax.adabelief(learning_rate=0.003) -opt_state=optimizer.init(params) - -def total_loss(params): - return jnp.linalg.norm(curvature_partial(params)+length_partial(params)+Baxis_average_partial(params))**2 - - -jit -def update(params,opt_state): - this_grad = grad(total_loss)(params) - updates, opt_state =optimizer.update(this_grad, opt_state) - params = optax.apply_updates(params, updates) - return params,opt_state - -for i in range(maximum_function_evaluations): - params,opt_state=update(params,opt_state) - if i % 3 == 0: - print('Objective function at iteration {:d}: {:.2E}'.format(i, total_loss(params))) - - -dofs_curves = jnp.reshape(params[:len_dofs_curves], (dofs_curves_shape)) -dofs_currents = params[len_dofs_curves:] -curves = Curves(dofs_curves, n_segments, nfp, stellsym) -new_coils = Coils(curves=curves, currents=dofs_currents*coils_initial.currents_scale) -params=new_coils.x -tracing_initial = Tracing(field=coils_initial, particles=particles, maxtime=t, model=model - ,times_to_trace=200,timestep=1.e-8,atol=1.e-5,rtol=1.e-5) -tracing_optimized = Tracing(field=new_coils, particles=particles, maxtime=t, model=model,times_to_trace=200,timestep=1.e-8,atol=1.e-5,rtol=1.e-5) - - -# Plot trajectories, before and after optimization -fig = plt.figure(figsize=(9, 8)) -ax1 = fig.add_subplot(221, projection='3d') -ax2 = fig.add_subplot(222, projection='3d') -ax3 = fig.add_subplot(223) -ax4 = fig.add_subplot(224) - -coils_initial.plot(ax=ax1, show=False) -tracing_initial.plot(ax=ax1, show=False) -for i, trajectory in enumerate(tracing_initial.trajectories): - ax3.plot(jnp.sqrt(trajectory[:,0]**2+trajectory[:,1]**2), trajectory[:, 2], label=f'Particle {i+1}') - -ax3.set_xlabel('R (m)') -ax3.set_ylabel('Z (m)') -#ax3.legend() -new_coils.plot(ax=ax2, show=False) -tracing_optimized.plot(ax=ax2, show=False) -for i, trajectory in enumerate(tracing_optimized.trajectories): - ax4.plot(jnp.sqrt(trajectory[:,0]**2+trajectory[:,1]**2), trajectory[:, 2], label=f'Particle {i+1}') -ax4.set_xlabel('R (m)') -ax4.set_ylabel('Z (m)')#ax4.legend() -plt.tight_layout() -# plt.savefig(f'opt_adam.pdf') -plt.savefig('optimize_coils_particle_confinement_guidingcenter_adam.png', dpi=300) -# # Save the coils to a json file -# coils_optimized.to_json("stellarator_coils.json") -# # Load the coils from a json file -# from essos.coils import Coils -# coils = Coils.from_json("stellarator_coils.json") - -# # Save results in vtk format to analyze in Paraview -# tracing_initial.to_vtk('trajectories_initial') -#tracing_optimized.to_vtk('trajectories_final') -#coils_initial.to_vtk('coils_initial') -#new_coils.to_vtk('coils_optimized') \ No newline at end of file diff --git a/examples/coil_optimization/optimize_coils_particle_confinement_guidingcenter_augmented_lagrangian.py b/examples/coil_optimization/optimize_coils_particle_confinement_guidingcenter_augmented_lagrangian.py deleted file mode 100644 index 21937a3..0000000 --- a/examples/coil_optimization/optimize_coils_particle_confinement_guidingcenter_augmented_lagrangian.py +++ /dev/null @@ -1,174 +0,0 @@ - -import os -number_of_processors_to_use = 1 # Parallelization, this should divide nparticles -os.environ["XLA_FLAGS"] = f'--xla_force_host_platform_device_count={number_of_processors_to_use}' -from time import time -import jax -print(jax.devices()) -jax.config.update("jax_enable_x64", True) -import jax.numpy as jnp -import matplotlib.pyplot as plt -from essos.surfaces import SurfaceRZFourier -from essos.dynamics import Particles, Tracing -from essos.coils import Coils, CreateEquallySpacedCurves,Curves -from essos.objective_functions import loss_particle_r_cross_max_constraint,loss_particle_gamma_c -from essos.objective_functions import loss_coil_curvature_new as loss_coil_curvature, loss_coil_length_new as loss_coil_length, loss_normB_axis_average,loss_Br,loss_iota -from functools import partial -import essos.augmented_lagrangian as alm - - - - - -# Optimization parameters -target_B_on_axis = 5.7 -max_coil_length = 31 -max_coil_curvature = 0.4 -nparticles = number_of_processors_to_use*10 -order_Fourier_series_coils = 4 -number_coil_points = 80 -maximum_function_evaluations = 1 -maxtimes = [1.e-5] -num_steps=100 -number_coils_per_half_field_period = 3 -number_of_field_periods = 2 -model = 'GuidingCenter' - -# Initialize coils -current_on_each_coil = 1.84e7 -major_radius_coils = 7.75 -minor_radius_coils = 4.45 -curves = CreateEquallySpacedCurves(n_curves=number_coils_per_half_field_period, - order=order_Fourier_series_coils, - R=major_radius_coils, r=minor_radius_coils, - n_segments=number_coil_points, - nfp=number_of_field_periods, stellsym=True) -coils_initial = Coils(curves=curves, currents=[current_on_each_coil]*number_coils_per_half_field_period) - - -len_dofs_curves = len(jnp.ravel(coils_initial.dofs_curves)) -nfp = coils_initial.nfp -stellsym = coils_initial.stellsym -n_segments = coils_initial.n_segments -dofs_curves_shape = coils_initial.dofs_curves.shape -currents_scale = coils_initial.currents_scale - - -# Initialize particles -phi_array = jnp.linspace(0, 2*jnp.pi, nparticles) -initial_xyz=jnp.array([major_radius_coils*jnp.cos(phi_array), major_radius_coils*jnp.sin(phi_array), 0*phi_array]).T -particles = Particles(initial_xyz=initial_xyz) - -t=maxtimes[0] -loss_partial = partial(loss_particle_gamma_c,particles=particles, dofs_curves=coils_initial.dofs_curves, currents_scale=currents_scale, nfp=nfp, n_segments=n_segments, stellsym=stellsym,maxtime=t,model=model,num_steps=num_steps) -curvature_partial=partial(loss_coil_curvature, dofs_curves=coils_initial.dofs_curves, currents_scale=currents_scale, nfp=nfp, n_segments=n_segments, stellsym=stellsym,max_coil_curvature=max_coil_curvature) -length_partial=partial(loss_coil_length, dofs_curves=coils_initial.dofs_curves, currents_scale=currents_scale, nfp=nfp, n_segments=n_segments, stellsym=stellsym,max_coil_length=max_coil_length) -Baxis_average_partial=partial(loss_normB_axis_average,dofs_curves=coils_initial.dofs_curves, currents_scale=currents_scale, nfp=nfp, n_segments=n_segments, stellsym=stellsym,npoints=15,target_B_on_axis=target_B_on_axis) -r_max_partial = partial(loss_particle_r_cross_max_constraint,target_r=0.4, particles=particles,dofs_curves=coils_initial.dofs_curves, currents_scale=currents_scale, nfp=nfp, n_segments=n_segments, stellsym=stellsym,maxtime=t,model=model,num_steps=num_steps) -iota_partial = partial(loss_iota,target_iota=0.5, particles=particles,dofs_curves=coils_initial.dofs_curves, currents_scale=currents_scale, nfp=nfp, n_segments=n_segments, stellsym=stellsym,maxtime=t,model=model,num_steps=num_steps) -Br_partial = partial(loss_Br, particles=particles,dofs_curves=coils_initial.dofs_curves, currents_scale=currents_scale, nfp=nfp, n_segments=n_segments, stellsym=stellsym,maxtime=t,model=model,num_steps=num_steps) - - -# Create the constraints -penalty = 1. #Intial penalty values -multiplier=0.5 #Initial lagrange multiplier values -sq_grad=0.0 #Initial square gradient parameter value for Mu adaptative -model_lagrangian='Standard' #Use standard augmented lagragian suitable for bounded optimizers -#Since we are using LBFGS-B from jaxopt, model_mu will be updated with tolerances so we do not need to difinte the model - -#Construct constraints -constraints = alm.combine( -alm.eq(curvature_partial, multiplier=multiplier,penalty=penalty,sq_grad=sq_grad), -alm.eq(length_partial, multiplier=multiplier,penalty=penalty,sq_grad=sq_grad), -alm.eq(Baxis_average_partial, multiplier=multiplier,penalty=penalty,sq_grad=sq_grad), -alm.eq(r_max_partial, multiplier=multiplier,penalty=penalty,sq_grad=sq_grad), -#alm.eq(Br_partial, multiplier=multiplier,penalty=penalty,sq_grad=sq_grad), -#alm.eq(iota_partial, multiplier=multiplier,penalty=penalty,sq_grad=sq_grad), -) - - - -beta=2. #penalty update parameter -mu_max=1.e4 #Maximum penalty parameter allowed -alpha=0.99 #These are parameters only used if gradient descent and adaaptative mu -gamma=1.e-2 -epsilon=1.e-8 -omega_tol=0.0001 #desired grad_tolerance, associated with grad of lagrangian to main parameters -eta_tol=0.001 #desired contraint tolerance, associated with variation of contraints - - - -#If loss=cost_function(x) is not prescribed, f(x)=0 is considered -ALM=alm.ALM_model_jaxopt_lbfgsb(constraints,loss=loss_partial,model_lagrangian=model_lagrangian,beta=beta,mu_max=mu_max,alpha=alpha,gamma=gamma,epsilon=epsilon,eta_tol=eta_tol,omega_tol=omega_tol) - -#Initializing lagrange multipliers -lagrange_params=constraints.init(coils_initial.x) -#parameters are a tuple of the primal/main optimisation parameters and the lagrange multipliers -params = coils_initial.x, lagrange_params -#This is just to initialize an empty state for the lagrange multiplier update and get some information -lag_state,grad,info=ALM.init(params) - -#Initializing first tolerances for the inner minimisation loop iteration -mu_average=alm.penalty_average(lagrange_params) -#omega=1.#1./mu_average -#eta=1000.#1./mu_average**0.1 -omega=1./mu_average -eta=1./mu_average**0.1 - -i=0 -while i<=maximum_function_evaluations and (jnp.linalg.norm(grad[0])>omega_tol or alm.norm_constraints(info[2])>eta_tol): - #One step of ALM optimization - params, lag_state,grad,info,eta,omega = ALM.update(params,lag_state,grad,info,eta,omega) - #if i % 5 == 0: - #print(f'i: {i}, loss f: {info[0]:g}, infeasibility: {alm.total_infeasibility(info[1]):g}') - print(f'i: {i}, loss f: {info[0]:g},loss L: {info[1]:g}, infeasibility: {alm.total_infeasibility(info[2]):g}') - #print('lagrange',params[1]) - i=i+1 - - -dofs_curves = jnp.reshape(params[0][:len_dofs_curves], (dofs_curves_shape)) -dofs_currents = params[0][len_dofs_curves:] -curves = Curves(dofs_curves, n_segments, nfp, stellsym) -new_coils = Coils(curves=curves, currents=dofs_currents*coils_initial.currents_scale) -params=new_coils.x -tracing_initial = Tracing(field=coils_initial, particles=particles, maxtime=t, model=model - ,times_to_trace=200,timestep=1.e-8,atol=1.e-5,rtol=1.e-5) -tracing_optimized = Tracing(field=new_coils, particles=particles, maxtime=t, model=model,times_to_trace=200,timestep=1.e-8,atol=1.e-5,rtol=1.e-5) - -#print('Final params',params) -#print(info[1]) -# Plot trajectories, before and after optimization -fig = plt.figure(figsize=(9, 8)) -ax1 = fig.add_subplot(221, projection='3d') -ax2 = fig.add_subplot(222, projection='3d') -ax3 = fig.add_subplot(223) -ax4 = fig.add_subplot(224) - -coils_initial.plot(ax=ax1, show=False) -tracing_initial.plot(ax=ax1, show=False) -for i, trajectory in enumerate(tracing_initial.trajectories): - ax3.plot(jnp.sqrt(trajectory[:,0]**2+trajectory[:,1]**2), trajectory[:, 2], label=f'Particle {i+1}') - -ax3.set_xlabel('R (m)') -ax3.set_ylabel('Z (m)') -#ax3.legend() -new_coils.plot(ax=ax2, show=False) -tracing_optimized.plot(ax=ax2, show=False) -for i, trajectory in enumerate(tracing_optimized.trajectories): - ax4.plot(jnp.sqrt(trajectory[:,0]**2+trajectory[:,1]**2), trajectory[:, 2], label=f'Particle {i+1}') -ax4.set_xlabel('R (m)') -ax4.set_ylabel('Z (m)')#ax4.legend() -plt.tight_layout() -plt.savefig(f'opt_constrained.pdf') - -# # Save the coils to a json file -# coils_optimized.to_json("stellarator_coils.json") -# # Load the coils from a json file -# from essos.coils import Coils -# coils = Coils.from_json("stellarator_coils.json") - -# # Save results in vtk format to analyze in Paraview -# tracing_initial.to_vtk('trajectories_initial') -#tracing_optimized.to_vtk('trajectories_final') -#coils_initial.to_vtk('coils_initial') -#new_coils.to_vtk('coils_optimized') diff --git a/examples/coil_optimization/optimize_coils_particle_confinement_guidingcenter_lbfgs.py b/examples/coil_optimization/optimize_coils_particle_confinement_guidingcenter_lbfgs.py deleted file mode 100644 index 20fad61..0000000 --- a/examples/coil_optimization/optimize_coils_particle_confinement_guidingcenter_lbfgs.py +++ /dev/null @@ -1,126 +0,0 @@ - -import os -number_of_processors_to_use = 1 # Parallelization, this should divide nparticles -os.environ["XLA_FLAGS"] = f'--xla_force_host_platform_device_count={number_of_processors_to_use}' -from jax import jit, value_and_grad -import jax.numpy as jnp -import matplotlib.pyplot as plt -from essos.dynamics import Particles, Tracing -from essos.coils import Coils, CreateEquallySpacedCurves,Curves -from essos.optimization import optimize_loss_function -from essos.objective_functions import loss_particle_r_cross_final,loss_particle_radial_drift,loss_particle_gamma_c -from essos.objective_functions import loss_coil_curvature_new,loss_coil_length_new,loss_normB_axis,loss_normB_axis_average -from functools import partial -import optax - - -# Optimization parameters -target_B_on_axis = 5.7 -max_coil_length = 31 -max_coil_curvature = 0.4 -nparticles = number_of_processors_to_use*10 -order_Fourier_series_coils = 4 -number_coil_points = 80 -maximum_function_evaluations = 3 -maxtimes = [2.e-5] -num_steps=100 -number_coils_per_half_field_period = 3 -number_of_field_periods = 2 -model = 'GuidingCenterAdaptative' - -# Initialize coils -current_on_each_coil = 1.84e7 -major_radius_coils = 7.75 -minor_radius_coils = 4.45 -curves = CreateEquallySpacedCurves(n_curves=number_coils_per_half_field_period, - order=order_Fourier_series_coils, - R=major_radius_coils, r=minor_radius_coils, - n_segments=number_coil_points, - nfp=number_of_field_periods, stellsym=True) -coils_initial = Coils(curves=curves, currents=[current_on_each_coil]*number_coils_per_half_field_period) - -len_dofs_curves = len(jnp.ravel(coils_initial.dofs_curves)) -nfp = coils_initial.nfp -stellsym = coils_initial.stellsym -n_segments = coils_initial.n_segments -dofs_curves_shape = coils_initial.dofs_curves.shape -currents_scale = coils_initial.currents_scale - -# Initialize particles -phi_array = jnp.linspace(0, 2*jnp.pi, nparticles) -initial_xyz=jnp.array([major_radius_coils*jnp.cos(phi_array), major_radius_coils*jnp.sin(phi_array), 0*phi_array]).T -particles = Particles(initial_xyz=initial_xyz) - -t=maxtimes[0] - -curvature_partial=partial(loss_coil_curvature_new, dofs_curves=coils_initial.dofs_curves, currents_scale=currents_scale, nfp=nfp, n_segments=n_segments, stellsym=stellsym,max_coil_curvature=max_coil_curvature) -length_partial=partial(loss_coil_length_new, dofs_curves=coils_initial.dofs_curves, currents_scale=currents_scale, nfp=nfp, n_segments=n_segments, stellsym=stellsym,max_coil_length=max_coil_length) -Baxis_average_partial=partial(loss_normB_axis_average,dofs_curves=coils_initial.dofs_curves, currents_scale=currents_scale, nfp=nfp, n_segments=n_segments, stellsym=stellsym,npoints=15,target_B_on_axis=target_B_on_axis) -r_max_partial = partial(loss_particle_r_cross_final, particles=particles,dofs_curves=coils_initial.dofs_curves, currents_scale=currents_scale, nfp=nfp, n_segments=n_segments, stellsym=stellsym,maxtime=t,model = model,num_steps=num_steps) -def total_loss(params): - return jnp.linalg.norm(jnp.concatenate((jnp.ravel(r_max_partial(params)),length_partial(params),curvature_partial(params),Baxis_average_partial(params))))**2 - -params=coils_initial.x -optimizer=optax.lbfgs() -opt_state=optimizer.init(params) - -@jit -def update(params,opt_state): - value, grad = value_and_grad(total_loss)(params) - updates, opt_state =optimizer.update(grad, opt_state, params, value=value, grad=grad, value_fn=total_loss) - params = optax.apply_updates(params, updates) - return params,opt_state - -for i in range(maximum_function_evaluations): - params,opt_state=update(params,opt_state) - if i % 3 == 0: - print('Objective function at iteration {:d}: {:.2E}'.format(i, total_loss(params))) - -dofs_curves = jnp.reshape(params[:len_dofs_curves], (dofs_curves_shape)) -dofs_currents = params[len_dofs_curves:] -curves = Curves(dofs_curves, n_segments, nfp, stellsym) -new_coils = Coils(curves=curves, currents=dofs_currents*coils_initial.currents_scale) -params=new_coils.x -tracing_initial = Tracing(field=coils_initial, particles=particles, maxtime=t, model=model - ,times_to_trace=200,timestep=1.e-8,atol=1.e-5,rtol=1.e-5) -tracing_optimized = Tracing(field=new_coils, particles=particles, maxtime=t, model=model,times_to_trace=200,timestep=1.e-8,atol=1.e-5,rtol=1.e-5) - - -# Plot trajectories, before and after optimization -fig = plt.figure(figsize=(9, 8)) -ax1 = fig.add_subplot(221, projection='3d') -ax2 = fig.add_subplot(222, projection='3d') -ax3 = fig.add_subplot(223) -ax4 = fig.add_subplot(224) - -coils_initial.plot(ax=ax1, show=False) -tracing_initial.plot(ax=ax1, show=False) -for i, trajectory in enumerate(tracing_initial.trajectories): - ax3.plot(jnp.sqrt(trajectory[:,0]**2+trajectory[:,1]**2), trajectory[:, 2], label=f'Particle {i+1}') - -ax3.set_xlabel('R (m)') -ax3.set_ylabel('Z (m)') -#ax3.legend() -new_coils.plot(ax=ax2, show=False) -tracing_optimized.plot(ax=ax2, show=False) -for i, trajectory in enumerate(tracing_optimized.trajectories): - ax4.plot(jnp.sqrt(trajectory[:,0]**2+trajectory[:,1]**2), trajectory[:, 2], label=f'Particle {i+1}') -ax4.set_xlabel('R (m)') -ax4.set_ylabel('Z (m)')#ax4.legend() -plt.tight_layout() -# plt.savefig(f'opt_lbfgs.pdf') -plt.show() - -# # Save the coils to a json file -# coils_optimized.to_json("stellarator_coils.json") -# # Load the coils from a json file -# from essos.coils import Coils -# coils = Coils.from_json("stellarator_coils.json") - -# # Save results in vtk format to analyze in Paraview -# tracing_initial.to_vtk('trajectories_initial') -#tracing_optimized.to_vtk('trajectories_final') -#coils_initial.to_vtk('coils_initial') -#new_coils.to_vtk('coils_optimized') - - diff --git a/examples/coil_optimization/optimize_coils_particle_confinement_loss_fraction_augmented_lagrangian.py b/examples/coil_optimization/optimize_coils_particle_confinement_loss_fraction_augmented_lagrangian.py deleted file mode 100644 index fb173f9..0000000 --- a/examples/coil_optimization/optimize_coils_particle_confinement_loss_fraction_augmented_lagrangian.py +++ /dev/null @@ -1,180 +0,0 @@ - -import os -number_of_processors_to_use = 8 # Parallelization, this should divide nparticles -os.environ["XLA_FLAGS"] = f'--xla_force_host_platform_device_count={number_of_processors_to_use}' -from time import time -import jax -print(jax.devices()) -jax.config.update("jax_enable_x64", True) -import jax.numpy as jnp -import matplotlib.pyplot as plt -from essos.surfaces import SurfaceRZFourier, SurfaceClassifier -from essos.dynamics import Particles, Tracing -from essos.coils import Coils, CreateEquallySpacedCurves,Curves -from essos.objective_functions import loss_lost_fraction,loss_lost_fraction_times -from essos.objective_functions import loss_coil_curvature,loss_coil_length,loss_normB_axis_average,loss_Br,loss_iota -from functools import partial -import essos.augmented_lagrangian as alm - - - - - -# Optimization parameters -target_B_on_axis = 5.7 -max_coil_length = 31 -max_coil_curvature = 0.4 -nparticles = number_of_processors_to_use*1 -order_Fourier_series_coils = 4 -number_coil_points = 80 -maximum_function_evaluations = 10 -maxtimes = [1.e-2] -timestep=1.e-8 -num_steps=100 -number_coils_per_half_field_period = 3 -number_of_field_periods = 2 -model = 'GuidingCenterAdaptative' - -# Initialize coils -current_on_each_coil = 1.84e7 -major_radius_coils = 7.75 -minor_radius_coils = 4.45 -curves = CreateEquallySpacedCurves(n_curves=number_coils_per_half_field_period, - order=order_Fourier_series_coils, - R=major_radius_coils, r=minor_radius_coils, - n_segments=number_coil_points, - nfp=number_of_field_periods, stellsym=True) -coils_initial = Coils(curves=curves, currents=[current_on_each_coil]*number_coils_per_half_field_period) - - -len_dofs_curves = len(jnp.ravel(coils_initial.dofs_curves)) -nfp = coils_initial.nfp -stellsym = coils_initial.stellsym -n_segments = coils_initial.n_segments -dofs_curves_shape = coils_initial.dofs_curves.shape -currents_scale = coils_initial.currents_scale - -ntheta=30 -nphi=30 -input = os.path.join(os.path.dirname(__file__), '..', 'input_files', 'input.toroidal_surface') -surface= SurfaceRZFourier.from_input_file(input, ntheta=ntheta, nphi=nphi, range_torus='full torus') -timeI=time() -boundary=SurfaceClassifier(surface,h=0.1) -print(f"ESSOS boundary took {time()-timeI:.2f} seconds") -#print('Final params',params) -#print(info[1]) -# Plot trajectories, before and after optimization - - -# Initialize particles -phi_array = jnp.linspace(0, 2*jnp.pi, nparticles) -initial_xyz=jnp.array([major_radius_coils*jnp.cos(phi_array), major_radius_coils*jnp.sin(phi_array), 0*phi_array]).T -particles = Particles(initial_xyz=initial_xyz) - -t=maxtimes[0] -loss_partial = partial(loss_lost_fraction,particles=particles, dofs_curves=coils_initial.dofs_curves, currents_scale=currents_scale, nfp=nfp, n_segments=n_segments, stellsym=stellsym,maxtime=t,timestep=timestep,model=model,num_steps=num_steps,boundary=boundary) -jax.grad(loss_partial)(coils_initial.x) - - -curvature_partial=partial(loss_coil_curvature, dofs_curves=coils_initial.dofs_curves, currents_scale=currents_scale, nfp=nfp, n_segments=n_segments, stellsym=stellsym,max_coil_curvature=max_coil_curvature) -length_partial=partial(loss_coil_length, dofs_curves=coils_initial.dofs_curves, currents_scale=currents_scale, nfp=nfp, n_segments=n_segments, stellsym=stellsym,max_coil_length=max_coil_length) -Baxis_average_partial=partial(loss_normB_axis_average,dofs_curves=coils_initial.dofs_curves, currents_scale=currents_scale, nfp=nfp, n_segments=n_segments, stellsym=stellsym,npoints=15,target_B_on_axis=target_B_on_axis) - -# Create the constraints -penalty = 1. #Intial penalty values -multiplier=0.5 #Initial lagrange multiplier values -sq_grad=0.0 #Initial square gradient parameter value for Mu adaptative -model_lagrangian='Standard' #Use standard augmented lagragian suitable for bounded optimizers -#Since we are using LBFGS-B from jaxopt, model_mu will be updated with tolerances so we do not need to difinte the model - -#Construct constraints -constraints = alm.combine( -alm.eq(curvature_partial, multiplier=multiplier,penalty=penalty,sq_grad=sq_grad), -alm.eq(length_partial, multiplier=multiplier,penalty=penalty,sq_grad=sq_grad), -alm.eq(Baxis_average_partial, multiplier=multiplier,penalty=penalty,sq_grad=sq_grad), -alm.eq(loss_partial, multiplier=multiplier,penalty=penalty,sq_grad=sq_grad), -) - - - -beta=2. #penalty update parameter -mu_max=1.e4 #Maximum penalty parameter allowed -alpha=0.99 #These are parameters only used if gradient descent and adaaptative mu -gamma=1.e-2 -epsilon=1.e-8 -omega_tol=0.0001 #desired grad_tolerance, associated with grad of lagrangian to main parameters -eta_tol=0.001 #desired contraint tolerance, associated with variation of contraints - - - -#If loss=cost_function(x) is not prescribed, f(x)=0 is considered -ALM=alm.ALM_model_jaxopt_lbfgsb(constraints,model_lagrangian=model_lagrangian,beta=beta,mu_max=mu_max,alpha=alpha,gamma=gamma,epsilon=epsilon,eta_tol=eta_tol,omega_tol=omega_tol) - -#Initializing lagrange multipliers -lagrange_params=constraints.init(coils_initial.x) -#parameters are a tuple of the primal/main optimisation parameters and the lagrange multipliers -params = coils_initial.x, lagrange_params -#This is just to initialize an empty state for the lagrange multiplier update and get some information -lag_state,grad,info=ALM.init(params) - -#Initializing first tolerances for the inner minimisation loop iteration -mu_average=alm.penalty_average(lagrange_params) -#omega=1.#1./mu_average -#eta=1000.#1./mu_average**0.1 -omega=1./mu_average -eta=1./mu_average**0.1 - -i=0 -while i<=maximum_function_evaluations and (jnp.linalg.norm(grad[0])>omega_tol or alm.norm_constraints(info[2])>eta_tol): - #One step of ALM optimization - params, lag_state,grad,info,eta,omega = ALM.update(params,lag_state,grad,info,eta,omega) - print(f'i: {i}, loss f: {info[0]:g},loss L: {info[1]:g}, infeasibility: {alm.total_infeasibility(info[2]):g}') - #print('lagrange',params[1]) - i=i+1 - - -dofs_curves = jnp.reshape(params[0][:len_dofs_curves], (dofs_curves_shape)) -dofs_currents = params[0][len_dofs_curves:] -curves = Curves(dofs_curves, n_segments, nfp, stellsym) -new_coils = Coils(curves=curves, currents=dofs_currents*coils_initial.currents_scale) -params=new_coils.x -tracing_initial = Tracing(field=coils_initial, particles=particles, maxtime=t, model=model,times_to_trace=num_steps,timestep=timestep,boundary=boundary) -tracing_optimized = Tracing(field=new_coils, particles=particles, maxtime=t, model=model,times_to_trace=num_steps,timestep=timestep,boundary=boundary) - -#print('Final params',params) -#print(info[1]) -# Plot trajectories, before and after optimization -fig = plt.figure(figsize=(9, 8)) -ax1 = fig.add_subplot(221, projection='3d') -ax2 = fig.add_subplot(222, projection='3d') -ax3 = fig.add_subplot(223) -ax4 = fig.add_subplot(224) - -coils_initial.plot(ax=ax1, show=False) -tracing_initial.plot(ax=ax1, show=False) -for i, trajectory in enumerate(tracing_initial.trajectories): - ax3.plot(jnp.sqrt(trajectory[:,0]**2+trajectory[:,1]**2), trajectory[:, 2], label=f'Particle {i+1}') - -ax3.set_xlabel('R (m)') -ax3.set_ylabel('Z (m)') -#ax3.legend() -new_coils.plot(ax=ax2, show=False) -tracing_optimized.plot(ax=ax2, show=False) -for i, trajectory in enumerate(tracing_optimized.trajectories): - ax4.plot(jnp.sqrt(trajectory[:,0]**2+trajectory[:,1]**2), trajectory[:, 2], label=f'Particle {i+1}') -ax4.set_xlabel('R (m)') -ax4.set_ylabel('Z (m)')#ax4.legend() -plt.tight_layout() -plt.show() - -# # Save the coils to a json file -# coils_optimized.to_json("stellarator_coils.json") -# # Load the coils from a json file -# from essos.coils import Coils -# coils = Coils.from_json("stellarator_coils.json") - -# # Save results in vtk format to analyze in Paraview -# tracing_initial.to_vtk('trajectories_initial') -#tracing_optimized.to_vtk('trajectories_final') -#coils_initial.to_vtk('coils_initial') -#new_coils.to_vtk('coils_optimized') diff --git a/examples/coil_optimization/optimize_coils_vmec_surface_augmented_lagrangian.py b/examples/coil_optimization/optimize_coils_vmec_surface_augmented_lagrangian.py deleted file mode 100644 index e675463..0000000 --- a/examples/coil_optimization/optimize_coils_vmec_surface_augmented_lagrangian.py +++ /dev/null @@ -1,189 +0,0 @@ -import os -number_of_processors_to_use = 1 # Parallelization, this should divide ntheta*nphi -os.environ["XLA_FLAGS"] = f'--xla_force_host_platform_device_count={number_of_processors_to_use}' -from time import time -import jax.numpy as jnp -import matplotlib.pyplot as plt -from essos.surfaces import BdotN_over_B -from essos.coils import Coils, CreateEquallySpacedCurves,Curves -from essos.fields import Vmec, BiotSavart -from essos.objective_functions import loss_BdotN_only_constraint,loss_coil_curvature_new,loss_coil_length_new,loss_BdotN_only -from essos.objective_functions import loss_coil_curvature_new as loss_coil_curvature, loss_coil_length_new as loss_coil_length -from essos.objective_functions import loss_BdotN -from essos.optimization import optimize_loss_function - -import essos.augmented_lagrangian as alm -from functools import partial - -# Optimization parameters -maximum_function_evaluations=10 -max_coil_length = 40 -max_coil_curvature = 0.5 -bdotn_tol=1.e-6 -order_Fourier_series_coils = 6 -number_coil_points = order_Fourier_series_coils*10 -number_coils_per_half_field_period = 4 -ntheta=32 -nphi=32 -#Tolerance for no normal (no ALM) optimization -tolerance_optimization = 1e-5 - -# Initialize VMEC field -vmec = Vmec(os.path.join(os.path.dirname(__file__), '..', 'input_files', - 'wout_LandremanPaul2021_QA_reactorScale_lowres.nc'), - ntheta=ntheta, nphi=nphi, range_torus='half period') - -# Initialize coils -current_on_each_coil = 1 -number_of_field_periods = vmec.nfp -major_radius_coils = vmec.r_axis -minor_radius_coils = vmec.r_axis/1.5 -curves = CreateEquallySpacedCurves(n_curves=number_coils_per_half_field_period, - order=order_Fourier_series_coils, - R=major_radius_coils, r=minor_radius_coils, - n_segments=number_coil_points, - nfp=number_of_field_periods, stellsym=True) -coils_initial = Coils(curves=curves, currents=[current_on_each_coil]*number_coils_per_half_field_period) - -len_dofs_curves = len(jnp.ravel(coils_initial.dofs_curves)) -nfp = coils_initial.nfp -stellsym = coils_initial.stellsym -n_segments = coils_initial.n_segments -dofs_curves = coils_initial.dofs_curves -currents_scale = coils_initial.currents_scale -dofs_curves_shape = coils_initial.dofs_curves.shape - - - - -# Create the constraints -penalty = 0.1 #Intial penalty values -multiplier=0.5 #Initial lagrange multiplier values -sq_grad=0.0 #Initial square gradient parameter value for Mu adaptative -model_lagrangian='Standard' #Use standard augmented lagragian suitable for bounded optimizers -#Since we are using LBFGS-B from jaxopt, model_mu will be updated with tolerances so we do not need to difinte the model - - -curvature_partial=partial(loss_coil_curvature, dofs_curves=coils_initial.dofs_curves, currents_scale=currents_scale, nfp=nfp, n_segments=n_segments, stellsym=stellsym,max_coil_curvature=max_coil_curvature) -length_partial=partial(loss_coil_length, dofs_curves=coils_initial.dofs_curves, currents_scale=currents_scale, nfp=nfp, n_segments=n_segments, stellsym=stellsym,max_coil_length=max_coil_length) -bdotn_partial=partial(loss_BdotN_only_constraint, vmec=vmec, dofs_curves=coils_initial.dofs_curves, currents_scale=currents_scale, nfp=nfp,n_segments=n_segments, stellsym=stellsym,target_tol=bdotn_tol) -bdotn_only_partial=partial(loss_BdotN_only, vmec=vmec, dofs_curves=coils_initial.dofs_curves, currents_scale=currents_scale, nfp=nfp,n_segments=n_segments, stellsym=stellsym) - -#Construct constraints -constraints = alm.combine( -alm.eq(curvature_partial,model_lagrangian=model_lagrangian, multiplier=multiplier,penalty=penalty,sq_grad=sq_grad), -alm.eq(length_partial,model_lagrangian=model_lagrangian, multiplier=multiplier,penalty=penalty,sq_grad=sq_grad), -alm.eq(bdotn_partial,model_lagrangian=model_lagrangian, multiplier=multiplier,penalty=penalty,sq_grad=sq_grad) -) - - - -beta=2. #penalty update parameter -mu_max=1.e4 #Maximum penalty parameter allowed -alpha=0.99 #These are parameters only used if gradient descent and adaaptative mu -gamma=1.e-2 -epsilon=1.e-8 -omega_tol=1.e-7 #desired grad_tolerance, associated with grad of lagrangian to main parameters -eta_tol=1.e-7 #desired contraint tolerance, associated with variation of contraints - - - -#If loss=cost_function(x) is not prescribed, f(x)=0 is considered -ALM=alm.ALM_model_jaxopt_lbfgsb(constraints,model_lagrangian=model_lagrangian,beta=beta,mu_max=mu_max,alpha=alpha,gamma=gamma,epsilon=epsilon,eta_tol=eta_tol,omega_tol=omega_tol) - -#Initializing lagrange multipliers -lagrange_params=constraints.init(coils_initial.x) -#parameters are a tuple of the primal/main optimisation parameters and the lagrange multipliers -params = coils_initial.x, lagrange_params -#This is just to initialize an empty state for the lagrange multiplier update and get some information -lag_state,grad,info=ALM.init(params) - -#Initializing first tolerances for the inner minimisation loop iteration -mu_average=alm.penalty_average(lagrange_params) -omega=1./mu_average -eta=1./mu_average**0.1 - - - - - -# Optimize coils -print(f'Optimizing coils with {maximum_function_evaluations} function evaluations no ALM.') -time0 = time() -coils_optimized = optimize_loss_function(loss_BdotN, initial_dofs=coils_initial.x, coils=coils_initial, tolerance_optimization=tolerance_optimization, - maximum_function_evaluations=maximum_function_evaluations, vmec=vmec, - max_coil_length=max_coil_length, max_coil_curvature=max_coil_curvature,) -print(f"Optimization took {time()-time0:.2f} seconds") - - - - - -# Optimize coils -print(f'Optimizing coils with {maximum_function_evaluations} function evaluations using ALM.') -time0 = time() - - -i=0 -while i<=maximum_function_evaluations and (jnp.linalg.norm(grad[0])>omega_tol or alm.norm_constraints(info[2])>eta_tol): - #One step of ALM optimization - params, lag_state,grad,info,eta,omega = ALM.update(params,lag_state,grad,info,eta,omega) - #if i % 5 == 0: - #print(f'i: {i}, loss f: {info[0]:g}, infeasibility: {alm.total_infeasibility(info[1]):g}') - print(f'i: {i}, loss f: {info[0]:g},loss L: {info[1]:g}, infeasibility: {alm.total_infeasibility(info[2]):g}') - #print('lagrange',params[1]) - i=i+1 - - - -dofs_curves = jnp.reshape(params[0][:len_dofs_curves], (dofs_curves_shape)) -dofs_currents = params[0][len_dofs_curves:] -curves = Curves(dofs_curves, n_segments, nfp, stellsym) -coils_optimized_alm = Coils(curves=curves, currents=dofs_currents*coils_initial.currents_scale) - -print(f"Optimization took {time()-time0:.2f} seconds") - - -BdotN_over_B_initial = BdotN_over_B(vmec.surface, BiotSavart(coils_initial)) -BdotN_over_B_optimized = BdotN_over_B(vmec.surface, BiotSavart(coils_optimized)) -curvature=jnp.mean(BiotSavart(coils_optimized).coils.curvature, axis=1) -length=jnp.max(jnp.ravel(BiotSavart(coils_optimized).coils.length)) -BdotN_over_B_optimized_alm = BdotN_over_B(vmec.surface, BiotSavart(coils_optimized_alm)) -curvature_alm=jnp.mean(BiotSavart(coils_optimized_alm).coils.curvature, axis=1) -length_alm=jnp.max(jnp.ravel(BiotSavart(coils_optimized_alm).coils.length)) - - -print(f"Maximum allowed curvature target: ",max_coil_curvature) -print(f"Maximum allowed length target: ",max_coil_length) -print(f"Mean curvature without ALM: ",curvature) -print(f"Length withou ALM:", length) -print(f"Mean curvature with ALM: ",curvature_alm) -print(f"Length with ALM:", length_alm) -print(f"Maximum BdotN/B before optimization: {jnp.max(BdotN_over_B_initial):.2e}") -print(f"Maximum BdotN/B after optimization without ALM: {jnp.max(BdotN_over_B_optimized):.2e}") -print(f"Maximum BdotN/B after optimization with ALM: {jnp.max(BdotN_over_B_optimized_alm):.2e}") -# Plot coils, before and after optimization -fig = plt.figure(figsize=(8, 4)) -ax1 = fig.add_subplot(121, projection='3d') -ax2 = fig.add_subplot(122, projection='3d') -coils_initial.plot(ax=ax1, show=False) -vmec.surface.plot(ax=ax1, show=False) -coils_optimized.plot(ax=ax2, show=False, label='Optimized no ALM') -coils_optimized_alm.plot(ax=ax2, show=False,color='orange', label='Optimized with ALM') -vmec.surface.plot(ax=ax2, show=False) -plt.legend() -plt.tight_layout() -plt.show() - -# # Save the coils to a json file -# coils_optimized.to_json("stellarator_coils.json") -# # Load the coils from a json file -# from essos.coils import Coils -# coils = Coils.from_json("stellarator_coils.json") - -# # Save results in vtk format to analyze in Paraview -# from essos.fields import BiotSavart -# vmec.surface.to_vtk('surface_initial', field=BiotSavart(coils_initial)) -# vmec.surface.to_vtk('surface_final', field=BiotSavart(coils_optimized)) -# coils_initial.to_vtk('coils_initial') -# coils_optimized.to_vtk('coils_optimized') \ No newline at end of file diff --git a/examples/coil_optimization/optimize_coils_vmec_surface_augmented_lagrangian_comparison.py b/examples/coil_optimization/optimize_coils_vmec_surface_augmented_lagrangian_comparison.py index 0f40631..72818bb 100644 --- a/examples/coil_optimization/optimize_coils_vmec_surface_augmented_lagrangian_comparison.py +++ b/examples/coil_optimization/optimize_coils_vmec_surface_augmented_lagrangian_comparison.py @@ -11,15 +11,14 @@ import essos.augmented_lagrangian as alm from functools import partial -# In this exmple, `scipy.optimize.least_squares` is used for the normal optimization, but any other optimizer, e.g. from +# In this exmple, `scipy.optimize.least_squares` is used for the normal optimization, but any other optimizer, e.g. from # `scipy.optimize.minimize` or `jaxopt`, can be used as well and may even be preferable. from scipy.optimize import least_squares # Optimization parameters maximum_function_evaluations=100 - -input_filepath = os.path.join(os.path.dirname(__name__), "input_files") +input_filepath = os.path.join(os.path.dirname(__file__), "..", "input_files") vmec_input = os.path.join(input_filepath, 'wout_LandremanPaul2021_QA_reactorScale_lowres.nc') surface = SurfaceRZFourier.from_wout_file(vmec_input, s=1, ntheta=32, nphi=32, range_torus='half period') @@ -45,14 +44,12 @@ EXPORT = False - init_curves = CreateEquallySpacedCurves(N_COILS, FOURIER_ORDER, LARGE_R, SMALL_R, n_segments=N_SEGMENTS, nfp=NFP, stellsym=STELLSYM) init_coils = Coils(curves=init_curves, currents=[COIL_CURRENT]*N_COILS) init_field = BiotSavart(init_coils) init_surface=surface - """ Creating the loss functions """ def loss(field, surface): return jnp.sum(jnp.abs(BdotN_over_B(surface, field))) @@ -100,7 +97,7 @@ def loss_curvature(field): omega=1./penalty eta=1./penalty**0.1 sq_grad=0.0 #Initial square gradient parameter value for Mu adaptative -model_lagrangian='Standard' #Use standard augmented lagragian suitable for bounded optimizers +model_lagrangian='Standard' #Use standard augmented lagragian suitable for bounded optimizers #Since we are using LBFGS-B from jaxopt, model_mu will be updated with tolerances so we do not need to difinte the model model_mu='Tolerance' @@ -133,7 +130,7 @@ def loss_curvature(field): C_normal_field_constraint.dependencies = {"field": init_field} C_length_constraint.dependencies = {"field": init_field} C_curvature_constraint.dependencies = {"field": init_field} -C_Total_constraint.dependencies = {"field": init_field} +C_Total_constraint.dependencies = {"field": init_field} #If loss=cost_function(x) is not prescribed, f(x)=0 is considered, uncomment second line to use B dot N as a loss and not a constraint @@ -153,7 +150,7 @@ def loss_curvature(field): i=0 while i<=maximum_function_evaluations and (jnp.linalg.norm(grad[0])>omega_tol or alm.norm_constraints(info[2])>eta_tol): #One step of ALM optimization - params, lag_state,grad,info = ALM.update(params,lag_state,grad,info) + params, lag_state,grad,info = ALM.update(params,lag_state,grad,info) #if i % 5 == 0: #print(f'i: {i}, loss f: {info[0]:g}, infeasibility: {alm.total_infeasibility(info[1]):g}') print(f'i: {i}, loss f: {info[0]:g},loss L: {info[1]:g}, infeasibility: {alm.total_infeasibility(info[2]):g}') @@ -177,7 +174,7 @@ def loss_curvature(field): t_end = time() print(f"\nOptimization took {t_end - t_start:.2f} seconds") -print("Initial loss:", L_total(L_total.starting_dofs)) +print("Initial loss:", L_total(L_total.starting_dofs)) print("Loss after optimization:", L_total(res.x)) opt_field = L_total.dofs_to_pytree(res.x)["field"] @@ -185,14 +182,14 @@ def loss_curvature(field): print(f"\nOptimization took {t_end - t_start:.2f} seconds") -print("Initial B dot N:", jnp.max(BdotN_over_B(surface, init_field))) +print("Initial B dot N:", jnp.max(BdotN_over_B(surface, init_field))) print("B dot N after optimization:", jnp.max(BdotN_over_B(surface, opt_field))) print("B dot N after optimization alm:", jnp.max(BdotN_over_B(surface, opt_field_alm))) -print("Initial curvature :", jnp.average(init_field.coils.curvature,axis=0)) +print("Initial curvature :", jnp.average(init_field.coils.curvature,axis=0)) print("Curvature after optimization:",jnp.average(opt_field.coils.curvature,axis=0)) print("Curvature after optimization alm:",jnp.average(opt_field_alm.coils.curvature,axis=0)) print("Curvature target:",CURVATURE_TARGET) -print("Initial length :", init_field.coils.length) +print("Initial length :", init_field.coils.length) print("Length after optimization:",opt_field.coils.length) print("Length after optimization alm:",opt_field_alm.coils.length) print("Length target:",LENGTH_TARGET) @@ -226,4 +223,4 @@ def loss_curvature(field): surface.to_vtk(os.path.join(output_filepath, "final_surface_vmec_surface.json"), field=opt_field) init_coils.to_vtk(os.path.join(output_filepath, "init_coils_vmec_surface.json")) opt_coils.to_vtk(os.path.join(output_filepath, "opt_coils_vmec_surface.json")) - opt_coils_alm.to_vtk(os.path.join(output_filepath, "opt_coils_alm_vmec_surface.json")) \ No newline at end of file + opt_coils_alm.to_vtk(os.path.join(output_filepath, "opt_coils_alm_vmec_surface.json")) diff --git a/examples/coil_optimization/optimize_coils_vmec_surface_augmented_lagrangian_stochastic.py b/examples/coil_optimization/optimize_coils_vmec_surface_augmented_lagrangian_stochastic.py index e7ee6e8..b289220 100644 --- a/examples/coil_optimization/optimize_coils_vmec_surface_augmented_lagrangian_stochastic.py +++ b/examples/coil_optimization/optimize_coils_vmec_surface_augmented_lagrangian_stochastic.py @@ -2,177 +2,181 @@ number_of_processors_to_use = 1 # Parallelization, this should divide ntheta*nphi os.environ["XLA_FLAGS"] = f'--xla_force_host_platform_device_count={number_of_processors_to_use}' from time import time + +import jax import jax.numpy as jnp import matplotlib.pyplot as plt -from essos.surfaces import BdotN_over_B -from essos.coils import Coils, CreateEquallySpacedCurves,Curves -from essos.fields import Vmec, BiotSavart -from essos.objective_functions import loss_BdotN_only_constraint_stochastic,loss_coil_curvature_new,loss_coil_length_new,loss_BdotN_only_stochastic -from essos.objective_functions import loss_coil_curvature_new as loss_coil_curvature, loss_coil_length_new as loss_coil_length -from essos.coil_perturbation import GaussianSampler import essos.augmented_lagrangian as alm -from functools import partial - -# Optimization parameters -maximum_function_evaluations=10 -max_coil_length = 40 -max_coil_curvature = 0.5 -bdotn_tol=1.e-6 -order_Fourier_series_coils = 6 -number_coil_points = order_Fourier_series_coils*10 -number_coils_per_half_field_period = 4 -ntheta=32 -nphi=32 +from essos.coil_perturbation import ( + GaussianSampler, + perturb_curves_statistic, + perturb_curves_systematic, +) +from essos.coils import Coils, CreateEquallySpacedCurves +from essos.fields import BiotSavart, Vmec +from essos.losses import custom_loss +from essos.surfaces import BdotN_over_B +""" Creating stochastic field losses """ +def copy_coils_from_field(field): + return field.coils.copy() +def perturbed_field_from_field(field, key, sampler): + coils = copy_coils_from_field(field) + base_key = jax.random.key(key) + split_keys = jax.random.split(base_key, 2) + coils = perturb_curves_systematic(coils, sampler, key=split_keys[0]) + coils = perturb_curves_statistic(coils, sampler, key=split_keys[1]) + return BiotSavart(coils) -# Initialize VMEC field -vmec = Vmec(os.path.join(os.path.dirname(__file__), '..', 'input_files', - 'wout_LandremanPaul2021_QA_reactorScale_lowres.nc'), - ntheta=ntheta, nphi=nphi, range_torus='full torus') -# Initialize coils -current_on_each_coil = 1 -number_of_field_periods = vmec.nfp -major_radius_coils = vmec.r_axis -minor_radius_coils = vmec.r_axis/1.5 -curves = CreateEquallySpacedCurves(n_curves=number_coils_per_half_field_period, - order=order_Fourier_series_coils, - R=major_radius_coils, r=minor_radius_coils, - n_segments=number_coil_points, - nfp=number_of_field_periods, stellsym=True) -coils_initial = Coils(curves=curves, currents=[current_on_each_coil]*number_coils_per_half_field_period) - -len_dofs_curves = len(jnp.ravel(coils_initial.dofs_curves)) -nfp = coils_initial.nfp -stellsym = coils_initial.stellsym -n_segments = coils_initial.n_segments -dofs_curves = coils_initial.dofs_curves -currents_scale = coils_initial.currents_scale -dofs_curves_shape = coils_initial.dofs_curves.shape - - - -#Sampling parameters -sigma=0.01 -length_scale=0.4*jnp.pi -n_derivs=2 -N_samples=10 #Number of samples for the stochastic perturbation -#Create a Gaussian sampler for perturbation -#This sampler will be used to perturb the coils -sampler=GaussianSampler(coils_initial.curves.quadpoints,sigma=sigma,length_scale=length_scale,n_derivs=n_derivs) - - - - -# Create the constraints -penalty = 0.1 #Intial penalty values -multiplier=0.5 #Initial lagrange multiplier values -sq_grad=0.0 #Initial square gradient parameter value for Mu adaptative -model_lagrangian='Standard' #Use standard augmented lagragian suitable for bounded optimizers -#Since we are using LBFGS-B from jaxopt, model_mu will be updated with tolerances so we do not need to difinte the model - - -curvature_partial=partial(loss_coil_curvature, dofs_curves=coils_initial.dofs_curves, currents_scale=currents_scale, nfp=nfp, n_segments=n_segments, stellsym=stellsym,max_coil_curvature=max_coil_curvature) -length_partial=partial(loss_coil_length, dofs_curves=coils_initial.dofs_curves, currents_scale=currents_scale, nfp=nfp, n_segments=n_segments, stellsym=stellsym,max_coil_length=max_coil_length) -bdotn_partial=partial(loss_BdotN_only_constraint_stochastic,sampler=sampler,N_samples=N_samples, vmec=vmec, dofs_curves=coils_initial.dofs_curves, currents_scale=currents_scale, nfp=nfp,n_segments=n_segments, stellsym=stellsym,target_tol=bdotn_tol) -bdotn_only_partial=partial(loss_BdotN_only_stochastic,sampler=sampler,N_samples=N_samples, vmec=vmec, dofs_curves=coils_initial.dofs_curves, currents_scale=currents_scale, nfp=nfp,n_segments=n_segments, stellsym=stellsym) - -#Construct constraints -constraints = alm.combine( -alm.eq(curvature_partial,model_lagrangian=model_lagrangian, multiplier=multiplier,penalty=penalty,sq_grad=sq_grad), -alm.eq(length_partial,model_lagrangian=model_lagrangian, multiplier=multiplier,penalty=penalty,sq_grad=sq_grad), -alm.eq(bdotn_partial,model_lagrangian=model_lagrangian, multiplier=multiplier,penalty=penalty,sq_grad=sq_grad) -) +def loss_bdotn_stochastic(field, surface, sampler, keys): + def perturbed_loss(key): + perturbed_field = perturbed_field_from_field(field, key, sampler) + bdotn_over_b = BdotN_over_B(surface, perturbed_field) + return jnp.sum(jnp.abs(bdotn_over_b)) + return jnp.mean(jax.vmap(perturbed_loss)(keys)) +def constraint_bdotn_stochastic(field, surface, sampler, keys, target_tol=1.0e-6): + def perturbed_square(key): + perturbed_field = perturbed_field_from_field(field, key, sampler) + return jnp.square(BdotN_over_B(surface, perturbed_field)) + expected_square = jnp.mean(jax.vmap(perturbed_square)(keys), axis=0) + return jnp.sqrt(jnp.sum(jnp.maximum(expected_square - target_tol, 0.0))) -beta=2. #penalty update parameter -mu_max=1.e4 #Maximum penalty parameter allowed -alpha=0.99 #These are parameters only used if gradient descent and adaaptative mu -gamma=1.e-2 -epsilon=1.e-8 -omega_tol=1.e-7 #desired grad_tolerance, associated with grad of lagrangian to main parameters -eta_tol=1.e-7 #desired contraint tolerance, associated with variation of contraints +def loss_length_constraint(field, max_coil_length): + return jnp.square(field.coils.length / max_coil_length - 1.0) -#If loss=cost_function(x) is not prescribed, f(x)=0 is considered -ALM=alm.ALM_model_jaxopt_lbfgsb(constraints,model_lagrangian=model_lagrangian,beta=beta,mu_max=mu_max,alpha=alpha,gamma=gamma,epsilon=epsilon,eta_tol=eta_tol,omega_tol=omega_tol) +def loss_curvature_constraint(field, max_coil_curvature): + pointwise_curvature_loss = jnp.square(jnp.maximum(field.coils.curvature - max_coil_curvature, 0.0)) + return jnp.mean(pointwise_curvature_loss * jnp.linalg.norm(field.coils.gamma_dash, axis=-1), axis=1) -#Initializing lagrange multipliers -lagrange_params=constraints.init(coils_initial.x) -#parameters are a tuple of the primal/main optimisation parameters and the lagrange multipliers -params = coils_initial.x, lagrange_params -#This is just to initialize an empty state for the lagrange multiplier update and get some information -lag_state,grad,info=ALM.init(params) -#Initializing first tolerances for the inner minimisation loop iteration -mu_average=alm.penalty_average(lagrange_params) -omega=1./mu_average -eta=1./mu_average**0.1 +# Optimization parameters +maximum_function_evaluations = 10 +MAX_COIL_LENGTH = 40.0 +MAX_COIL_CURVATURE = 0.5 +BDOTN_TARGET_TOL = 1.0e-6 +FOURIER_ORDER = 6 +N_SEGMENTS = FOURIER_ORDER * 10 +N_COILS = 4 +NTHETA = 32 +NPHI = 32 +input_filepath = os.path.join(os.path.dirname(__file__), "..", "input_files") +vmec_input = os.path.join(input_filepath, "wout_LandremanPaul2021_QA_reactorScale_lowres.nc") +""" Creating starting coils and surface """ +vmec = Vmec(vmec_input, ntheta=NTHETA, nphi=NPHI, range_torus="full torus") +surface = vmec.surface -# Optimize coils -print(f'Optimizing coils with {maximum_function_evaluations} function evaluations using stochastic and ALM.') +COIL_CURRENT = 1.0 +number_of_field_periods = vmec.nfp +major_radius_coils = vmec.r_axis +minor_radius_coils = vmec.r_axis / 1.5 +curves = CreateEquallySpacedCurves(n_curves=N_COILS,order=FOURIER_ORDER, R=major_radius_coils, + r=minor_radius_coils,n_segments=N_SEGMENTS, nfp=number_of_field_periods,stellsym=True) +coils_initial = Coils(curves=curves,currents=[COIL_CURRENT] * N_COILS) +field_initial = BiotSavart(coils_initial) + +""" Setting the stochastic sampling parameters """ +SIGMA = 0.01 +LENGTH_SCALE = 0.4 * jnp.pi +N_DERIVS = 2 +N_samples = 10 +sampler = GaussianSampler(coils_initial.curves.quadpoints, sigma=SIGMA, length_scale=LENGTH_SCALE, n_derivs=N_DERIVS) +stochastic_keys = jnp.arange(N_samples) + +""" Defining custom losses """ +L_normal_field = custom_loss(loss_bdotn_stochastic, "field", surface=surface, sampler=sampler, keys=stochastic_keys) +L_normal_field_constraint = custom_loss(constraint_bdotn_stochastic, "field", surface=surface, sampler=sampler, keys=stochastic_keys, target_tol=BDOTN_TARGET_TOL) +L_length_constraint = custom_loss(loss_length_constraint, "field", max_coil_length=MAX_COIL_LENGTH) +L_curvature_constraint = custom_loss(loss_curvature_constraint, "field", max_coil_curvature=MAX_COIL_CURVATURE) + +""" Defining total loss + setting dependencies """ +L_normal_field.dependencies = {"field": field_initial} +L_normal_field_constraint.dependencies = {"field": field_initial} +L_length_constraint.dependencies = {"field": field_initial} +L_curvature_constraint.dependencies = {"field": field_initial} + +""" Creating the constraints """ +penalty = 0.1 +multiplier = 0.5 +sq_grad = 0.0 +model_lagrangian = "Standard" + +beta = 2.0 +mu_max = 1.0e4 +alpha = 0.99 +gamma = 1.0e-2 +epsilon = 1.0e-8 +omega_tol = 1.0e-7 +eta_tol = 1.0e-7 + +normal_field_constraint = alm.eq(constraint_bdotn_stochastic, model_lagrangian=model_lagrangian, multiplier=multiplier, penalty=penalty, sq_grad=sq_grad) +length_constraint = alm.eq(loss_length_constraint, model_lagrangian=model_lagrangian, multiplier=multiplier, penalty=penalty, sq_grad=sq_grad) +curvature_constraint = alm.eq(loss_curvature_constraint, model_lagrangian=model_lagrangian, multiplier=multiplier, penalty=penalty, sq_grad=sq_grad) + +C_normal_field_constraint = alm.SelectiveConstraint(normal_field_constraint, "field", surface=surface, sampler=sampler, keys=stochastic_keys, target_tol=BDOTN_TARGET_TOL) +C_length_constraint = alm.SelectiveConstraint(length_constraint, "field", max_coil_length=MAX_COIL_LENGTH) +C_curvature_constraint = alm.SelectiveConstraint(curvature_constraint, "field", max_coil_curvature=MAX_COIL_CURVATURE) +C_total_constraint = alm.combine(C_normal_field_constraint, C_length_constraint, C_curvature_constraint) + +C_normal_field_constraint.dependencies = {"field": field_initial} +C_length_constraint.dependencies = {"field": field_initial} +C_curvature_constraint.dependencies = {"field": field_initial} +C_total_constraint.dependencies = {"field": field_initial} + +ALM = alm.ALM_model_jaxopt_lbfgsb(constraints=C_total_constraint, model_lagrangian=model_lagrangian, beta=beta, mu_max=mu_max, alpha=alpha, gamma=gamma, epsilon=epsilon, eta_tol=eta_tol, omega_tol=omega_tol) + +""" Optimizing with alm """ +lagrange_params = C_total_constraint.init(field_initial.dofs) +params = field_initial.dofs, lagrange_params +lag_state, grad, info = ALM.init(params) + +print(f"Optimizing coils with {maximum_function_evaluations} function evaluations using stochastic ALM.") time0 = time() +i = 0 +while i <= maximum_function_evaluations and (jnp.linalg.norm(grad[0]) > omega_tol or alm.norm_constraints(info[2]) > eta_tol): + params, lag_state, grad, info = ALM.update(params, lag_state, grad, info) + print(f"i: {i}, loss f: {info[0]:g}, loss L: {info[1]:g}, " f"infeasibility: {alm.total_infeasibility(info[2]):g}") + i += 1 -i=0 -while i<=maximum_function_evaluations and (jnp.linalg.norm(grad[0])>omega_tol or alm.norm_constraints(info[2])>eta_tol): - #One step of ALM optimization - params, lag_state,grad,info,eta,omega = ALM.update(params,lag_state,grad,info,eta,omega) - #if i % 5 == 0: - #print(f'i: {i}, loss f: {info[0]:g}, infeasibility: {alm.total_infeasibility(info[1]):g}') - print(f'i: {i}, loss f: {info[0]:g},loss L: {info[1]:g}, infeasibility: {alm.total_infeasibility(info[2]):g}') - #print('lagrange',params[1]) - i=i+1 - +field_optimized = C_normal_field_constraint.dofs_to_pytree(params[0])[0] +coils_optimized = field_optimized.coils +print(f"Stochastic optimization with ALM took {time() - time0:.2f} seconds") -dofs_curves = jnp.reshape(params[0][:len_dofs_curves], (dofs_curves_shape)) -dofs_currents = params[0][len_dofs_curves:] -curves = Curves(dofs_curves, n_segments, nfp, stellsym) -coils_optimized = Coils(curves=curves, currents=dofs_currents*coils_initial.currents_scale) +BdotN_over_B_initial = BdotN_over_B(surface, field_initial) +BdotN_over_B_optimized = BdotN_over_B(surface, field_optimized) +curvature = jnp.mean(field_optimized.coils.curvature, axis=1) +length = jnp.max(jnp.ravel(field_optimized.coils.length)) +stochastic_loss_initial = L_normal_field(L_normal_field.starting_dofs) +stochastic_loss_final = L_normal_field(params[0]) -print(f"Stochastic optimization with ALM took {time()-time0:.2f} seconds") - - -BdotN_over_B_initial = BdotN_over_B(vmec.surface, BiotSavart(coils_initial)) -BdotN_over_B_optimized = BdotN_over_B(vmec.surface, BiotSavart(coils_optimized)) -curvature=jnp.mean(BiotSavart(coils_optimized).coils.curvature, axis=1) -length=jnp.max(jnp.ravel(BiotSavart(coils_optimized).coils.length)) -print(f"Mean curvature: ",curvature) -print(f"Length:", length) +print("Mean curvature:", curvature) +print("Length:", length) +print(f"Stochastic |BdotN/B| loss before optimization: {stochastic_loss_initial:.2e}") +print(f"Stochastic |BdotN/B| loss after optimization: {stochastic_loss_final:.2e}") print(f"Maximum BdotN/B before optimization: {jnp.max(BdotN_over_B_initial):.2e}") print(f"Maximum BdotN/B after optimization: {jnp.max(BdotN_over_B_optimized):.2e}") -print(f"Average BdotN/B before optimization: {jnp.average(jnp.absolute(BdotN_over_B_initial)):.2e}") -print(f"Average BdotN/B after optimization: {jnp.average(jnp.absolute(BdotN_over_B_optimized)):.2e}") -# Plot coils, before and after optimization +print(f"Average BdotN/B before optimization: {jnp.average(jnp.abs(BdotN_over_B_initial)):.2e}") +print(f"Average BdotN/B after optimization: {jnp.average(jnp.abs(BdotN_over_B_optimized)):.2e}") + fig = plt.figure(figsize=(8, 4)) -ax1 = fig.add_subplot(121, projection='3d') -ax2 = fig.add_subplot(122, projection='3d') +ax1 = fig.add_subplot(121, projection="3d") +ax2 = fig.add_subplot(122, projection="3d") coils_initial.plot(ax=ax1, show=False) -vmec.surface.plot(ax=ax1, show=False) +surface.plot(ax=ax1, show=False) coils_optimized.plot(ax=ax2, show=False) -vmec.surface.plot(ax=ax2, show=False) +surface.plot(ax=ax2, show=False) plt.tight_layout() plt.show() - -# # Save the coils to a json file -# coils_optimized.to_json("stellarator_coils.json") -# # Load the coils from a json file -# from essos.coils import Coils -# coils = Coils.from_json("stellarator_coils.json") - -# # Save results in vtk format to analyze in Paraview -# from essos.fields import BiotSavart -# vmec.surface.to_vtk('surface_initial', field=BiotSavart(coils_initial)) -# vmec.surface.to_vtk('surface_final', field=BiotSavart(coils_optimized)) -# coils_initial.to_vtk('coils_initial') -# coils_optimized.to_vtk('coils_optimized') \ No newline at end of file diff --git a/examples/optimize_surface_quasisymmetry.py b/examples/optimize_surface_quasisymmetry.py deleted file mode 100644 index 1e4e259..0000000 --- a/examples/optimize_surface_quasisymmetry.py +++ /dev/null @@ -1,164 +0,0 @@ -import os -number_of_processors_to_use = 1 # Sharded arrays here have sizes 13 and 30; no single count >1 divides both, so this runs unparallelized. -os.environ["XLA_FLAGS"] = f'--xla_force_host_platform_device_count={number_of_processors_to_use}' -from essos.fields import BiotSavart, near_axis -from essos.dynamics import Particles, Tracing -from essos.surfaces import BdotN_over_B, SurfaceRZFourier, B_on_surface -from essos.coils import Coils, CreateEquallySpacedCurves, Curves -from essos.optimization import optimize_loss_function, new_nearaxis_from_x_and_old_nearaxis -from essos.objective_functions import (loss_coil_curvature, difference_B_gradB_onaxis, - loss_coil_length, loss_BdotN) -import jax.numpy as jnp -from functools import partial -from jax import jit, vmap, devices, device_put, grad, debug -from jax.sharding import Mesh, NamedSharding, PartitionSpec -from time import time -import matplotlib.pyplot as plt -from simsopt.mhd import Vmec, QuasisymmetryRatioResidual - -mesh = Mesh(devices(), ("dev",)) -sharding = NamedSharding(mesh, PartitionSpec("dev", None)) - -ntheta=30 -nphi=30 -input = os.path.join(os.path.dirname(__file__), 'input_files', 'input.rotating_ellipse') -surface_initial = SurfaceRZFourier.from_input_file(input, ntheta=ntheta, nphi=nphi, range_torus='half period') - -# Optimization parameters -max_coil_length = 50 -max_coil_curvature = 0.4 -order_Fourier_series_coils = 12 -number_coil_points = 80#order_Fourier_series_coils*10 -maximum_function_evaluations = 600 -number_coils_per_half_field_period = 4 -tolerance_optimization = 1e-8 - -# Initialize coils -current_on_each_coil = 1.714e7 -number_of_field_periods = surface_initial.nfp -major_radius_coils = surface_initial.dofs[0] -minor_radius_coils = major_radius_coils/1.3 -curves = CreateEquallySpacedCurves(n_curves=number_coils_per_half_field_period, - order=order_Fourier_series_coils, - R=major_radius_coils, r=minor_radius_coils, - n_segments=number_coil_points, - nfp=number_of_field_periods, stellsym=True) -coils_initial = Coils(curves=curves, currents=[current_on_each_coil]*number_coils_per_half_field_period) -coils_initial = optimize_loss_function(loss_BdotN, initial_dofs=coils_initial.x, coils=coils_initial, tolerance_optimization=tolerance_optimization, - maximum_function_evaluations=maximum_function_evaluations, surface=surface_initial, - max_coil_length=max_coil_length, max_coil_curvature=max_coil_curvature,) - -def B_on_surface(surface, field): - ntheta = surface.ntheta - nphi = surface.nphi - gamma = surface.gamma - gamma_reshaped = gamma.reshape(nphi * ntheta, 3) - gamma_sharded = device_put(gamma_reshaped, sharding) - B_on_surface = jit(vmap(field.B), in_shardings=sharding, out_shardings=sharding)(gamma_sharded) - B_on_surface = B_on_surface.reshape(nphi, ntheta, 3) - return B_on_surface - -# @partial(jit, static_argnames=['surface','field']) -def grad_AbsB_on_surface(surface, field): - ntheta = surface.ntheta - nphi = surface.nphi - gamma = surface.gamma - gamma_reshaped = gamma.reshape(nphi * ntheta, 3) - gamma_sharded = device_put(gamma_reshaped, sharding) - dAbsB_by_dX_on_surface = jit(vmap(field.dAbsB_by_dX), in_shardings=sharding, out_shardings=sharding)(gamma_sharded) - dAbsB_by_dX_on_surface = dAbsB_by_dX_on_surface.reshape(nphi, ntheta, 3) - return dAbsB_by_dX_on_surface - -# @partial(jit, static_argnames=['field']) -def B_dot_GradAbsB(points, field): - B = field.B(points) - GradAbsB = field.dAbsB_by_dX(points) - B_dot_GradAbsB = jnp.sum(B * GradAbsB, axis=-1) - return B_dot_GradAbsB - -# @partial(jit, static_argnames=['field']) -def grad_B_dot_GradAbsB(points, field): - return grad(B_dot_GradAbsB, argnums=0)(points, field) - -# @partial(jit, static_argnames=['surface','field']) -def grad_B_dot_GradAbsB_on_surface(surface, field): - ntheta = surface.ntheta - nphi = surface.nphi - gamma = surface.gamma - gamma_reshaped = gamma.reshape(nphi * ntheta, 3) - gamma_sharded = device_put(gamma_reshaped, sharding) - partial_grad_B_dot_GradAbsB = partial(grad_B_dot_GradAbsB, field=field) - grad_B_dot_GradAbsB_on_surface = jit(vmap(partial_grad_B_dot_GradAbsB), in_shardings=sharding, out_shardings=sharding)(gamma_sharded) - grad_B_dot_GradAbsB_on_surface = grad_B_dot_GradAbsB_on_surface.reshape(nphi, ntheta, 3) - return grad_B_dot_GradAbsB_on_surface - -# @partial(jit, static_argnames=['surface','field']) -def loss_normal_cross_GradB_dot_grad_B_dot_GradB_surface(surface, field): - gradAbsB_surface = grad_AbsB_on_surface(surface, field) - B_surface = B_on_surface(surface, field) - grad_B_dot_GradB_surface = grad_B_dot_GradAbsB_on_surface(surface, field) - normal_cross_GradB_surface = jnp.cross(surface.normal, gradAbsB_surface, axisa=-1, axisb=-1) - normal_cross_GradB_dot_grad_B_dot_GradB_surface = jnp.sum(normal_cross_GradB_surface * grad_B_dot_GradB_surface, axis=-1) - B_cross_GradB = jnp.cross(B_surface, gradAbsB_surface, axisa=-1, axisb=-1) - # B_cross_GradB_dot_grad_B_dot_GradB_surface = jnp.sum(B_cross_GradB * grad_B_dot_GradB_surface, axis=-1) - # debug.print("normal_cross_GradB_dot_grad_B_dot_GradB_surface: {}", jnp.sum(jnp.abs(normal_cross_GradB_dot_grad_B_dot_GradB_surface))) - # debug.print("B_cross_GradB_dot_grad_B_dot_GradB_surface: {}", jnp.sum(jnp.abs(B_cross_GradB_dot_grad_B_dot_GradB_surface))) - return normal_cross_GradB_dot_grad_B_dot_GradB_surface#, normal_cross_GradB_dot_grad_B_dot_GradB_surface + B_cross_GradB_dot_grad_B_dot_GradB_surface - -def vmec_qs_from_surface(filename): - vmec = Vmec(filename, verbose=False) - qs = QuasisymmetryRatioResidual(vmec, surfaces=[1], helicity_m=1, helicity_n=0) - return jnp.sum(jnp.abs(qs.residuals())) - -# @partial(jit, static_argnames=['surface_initial', 'n_segments']) -def qs_loss(surface_dofs, dofs_curves, dofs_currents, surface_initial, currents_scale=1, n_segments=100): - surface = SurfaceRZFourier(rc=surface_initial.rc, zs=surface_initial.zs, nfp=surface_initial.nfp, range_torus=surface_initial.range_torus, nphi=surface_initial.nphi, ntheta=surface_initial.ntheta) - surface.dofs = surface_dofs - curves = Curves(dofs_curves, n_segments, surface_initial.nfp) - coils = Coils(curves=curves, currents=dofs_currents*currents_scale) - # print("##############################") - print(f" initial max(BdotN/B): {jnp.max(BdotN_over_B(surface, BiotSavart(coils))):.2e}") - field1 = BiotSavart(coils) - coils = optimize_loss_function(loss_BdotN, initial_dofs=coils.x, coils=coils, tolerance_optimization=tolerance_optimization, - maximum_function_evaluations=maximum_function_evaluations, surface=surface, - max_coil_length=max_coil_length, max_coil_curvature=max_coil_curvature,disp=False) - field2 = BiotSavart(coils) - print(f" final max(BdotN/B): {jnp.max(BdotN_over_B(surface, field2)):.2e}") - loss1 = jnp.sum(jnp.abs(loss_normal_cross_GradB_dot_grad_B_dot_GradB_surface(surface, field1))) - loss2 = jnp.sum(jnp.abs(loss_normal_cross_GradB_dot_grad_B_dot_GradB_surface(surface, field2))) - return loss1, loss2, coils - -print(f'############################################') -dofs_old = surface_initial.dofs -new_dof_array = jnp.linspace(-1.0, 1.0, 10) -qs_ESSOS_loss1_array = [] -qs_ESSOS_loss2_array = [] -qs_VMEC_loss_array = [] -for dof in new_dof_array: - coils = coils_initial - print(f'dof: {dof}') - dofs = dofs_old.at[2].set(dof) - loss1, loss2, coils = qs_loss(dofs, coils.dofs_curves, coils.dofs_currents, surface_initial, currents_scale=coils.currents_scale, n_segments=number_coil_points) - qs_ESSOS_loss1_array.append(loss1) - qs_ESSOS_loss2_array.append(loss2) - - filename = 'input.rotating_ellipse_dof' - new_surface = SurfaceRZFourier(rc=surface_initial.rc, zs=surface_initial.zs, nfp=surface_initial.nfp, range_torus=surface_initial.range_torus, nphi=surface_initial.nphi, ntheta=surface_initial.ntheta) - new_surface.dofs = dofs - new_surface.to_vmec(filename) - new_surface.to_vtk('surface_dof') - qs = vmec_qs_from_surface(filename) - qs_VMEC_loss_array.append(qs) - print('VMEC qs:', qs) - print('ESSOS loss1:', loss1) - print('ESSOS loss2:', loss2) -qs_ESSOS_loss1_array = jnp.array(qs_ESSOS_loss1_array) -qs_ESSOS_loss2_array = jnp.array(qs_ESSOS_loss2_array) -qs_VMEC_loss_array = jnp.array(qs_VMEC_loss_array) -plt.plot(new_dof_array, qs_ESSOS_loss1_array/jnp.max(qs_ESSOS_loss1_array), label='ESSOS without coil optimization') -plt.plot(new_dof_array, qs_ESSOS_loss2_array/jnp.max(qs_ESSOS_loss2_array), label='ESSOS with coil optimization') -plt.plot(new_dof_array, qs_VMEC_loss_array/jnp.max(qs_VMEC_loss_array), label='VMEC') -plt.legend() -plt.xlabel('Dof') -plt.ylabel('Loss') -plt.show() diff --git a/examples/particle_tracing/trace_particles_coils_guidingcenter.py b/examples/particle_tracing/trace_particles_coils_guidingcenter.py index 55bbd90..b7a52a5 100644 --- a/examples/particle_tracing/trace_particles_coils_guidingcenter.py +++ b/examples/particle_tracing/trace_particles_coils_guidingcenter.py @@ -21,7 +21,7 @@ energy=4000*ONE_EV # Load coils and field -json_file = os.path.join(os.path.dirname(__file__), '../input_files', 'ESSOS_biot_savart_LandremanPaulQA.json') +json_file = os.path.join(os.path.dirname(__file__), '..', 'input_files', 'ESSOS_biot_savart_LandremanPaulQA.json') coils = Coils.from_json(json_file) field = BiotSavart(coils) diff --git a/examples/particle_tracing/trace_particles_coils_guidingcenter_with_classifier.py b/examples/particle_tracing/trace_particles_coils_guidingcenter_with_classifier.py index 6e12c42..aa05ff7 100644 --- a/examples/particle_tracing/trace_particles_coils_guidingcenter_with_classifier.py +++ b/examples/particle_tracing/trace_particles_coils_guidingcenter_with_classifier.py @@ -1,6 +1,8 @@ import os number_of_processors_to_use = 1 # Parallelization, this should divide nparticles os.environ["XLA_FLAGS"] = f'--xla_force_host_platform_device_count={number_of_processors_to_use}' +import jax +print(jax.devices()) from time import time from jax import block_until_ready import jax.numpy as jnp @@ -10,6 +12,7 @@ from essos.coils import Coils from essos.constants import ALPHA_PARTICLE_MASS, ALPHA_PARTICLE_CHARGE, FUSION_ALPHA_PARTICLE_ENERGY,ONE_EV from essos.dynamics import Tracing, Particles +from essos.objective_functions import normB_axis # Input parameters tmax = 1.e-4 @@ -22,16 +25,25 @@ rtol=1.e-7 energy=FUSION_ALPHA_PARTICLE_ENERGY - - -# Load coils and field +# # Load coils and field json_file = os.path.join(os.path.dirname(__file__), '..', 'input_files', 'QH_simple_scaled.json') -coils = Coils.from_simsopt(json_file) +coils = Coils.from_simsopt(json_file, nfp=4) field = BiotSavart(coils) +# json_file = os.path.join(os.path.dirname(__file__), '..', 'input_files', 'ESSOS_biot_savart_LandremanPaulQA.json') +# coils = Coils.from_json(json_file) +# field = BiotSavart(coils) - +#renormalize coisl to have B_target=5.7 on axis +B_axis_old=normB_axis(field,npoints=200) +#print(jnp.average(B_axis_old)) +B_target=5.7 +coils.dofs_currents=coils.dofs_currents*B_target/jnp.average(B_axis_old) +field=BiotSavart(coils) +#B_axis_new=normB_axis(field,npoints=200) +#print(jnp.average(B_axis_new)) # Load coils and field wout_file = os.path.join(os.path.dirname(__file__), '..', 'input_files','wout_QH_simple_scaled.nc') +# wout_file = os.path.join(os.path.dirname(__file__), "../input_files", "wout_LandremanPaul2021_QA_reactorScale_lowres.nc") vmec = Vmec(wout_file) timeI=time() @@ -46,8 +58,8 @@ print(f"Initialization performed") # Trace in ESSOS time0 = time() -tracing = block_until_ready(Tracing(field=field, model='GuidingCenterAdaptative', particles=particles, - maxtime=tmax, timestep=timestep,times_to_trace=times_to_trace, atol=atol,rtol=rtol,boundary=boundary)) +tracing = Tracing(field=field, model='GuidingCenterAdaptative', particles=particles, + maxtime=tmax, timestep=timestep,times_to_trace=times_to_trace, atol=atol,rtol=rtol,boundary=boundary) print(f"ESSOS tracing took {time()-time0:.2f} seconds") print(f"Final loss fraction: {tracing.loss_fractions[-1]*100:.2f}%") trajectories = tracing.trajectories @@ -64,7 +76,7 @@ tracing.plot(ax=ax1, show=False, n_trajectories_plot=nparticles) for i, trajectory in enumerate(trajectories): - ax2.plot(tracing.times, jnp.abs(tracing.energy[i]-particles.energy)/particles.energy, label=f'Particle {i+1}') + ax2.plot(tracing.times, jnp.abs(tracing.energy()[i]-particles.energy)/particles.energy, label=f'Particle {i+1}') ax3.plot(tracing.times, trajectory[:, 3]/particles.total_speed, label=f'Particle {i+1}') #ax4.plot(jnp.sqrt(trajectory[:,0]**2+trajectory[:,1]**2), trajectory[:, 2], label=f'Particle {i+1}') ax4.plot(jnp.sqrt(trajectory[:,0]**2+trajectory[:,1]**2), trajectory[:, 2], label=f'Particle {i+1}') @@ -81,7 +93,6 @@ plt.tight_layout() plt.show() - ## Save results in vtk format to analyze in Paraview # tracing.to_vtk('trajectories') # coils.to_vtk('coils') diff --git a/examples/particle_tracing/trace_particles_coils_guidingcenter_with_classifier_scaled_currents.py b/examples/particle_tracing/trace_particles_coils_guidingcenter_with_classifier_scaled_currents.py index 9d4e770..fd9c175 100644 --- a/examples/particle_tracing/trace_particles_coils_guidingcenter_with_classifier_scaled_currents.py +++ b/examples/particle_tracing/trace_particles_coils_guidingcenter_with_classifier_scaled_currents.py @@ -25,8 +25,6 @@ rtol=1.e-7 energy=FUSION_ALPHA_PARTICLE_ENERGY - - # Load coils and field json_file = os.path.join(os.path.dirname(__file__), '..', 'input_files', 'QH_simple_scaled.json') coils = Coils.from_simsopt(json_file,nfp=4) diff --git a/examples/particle_tracing/trace_particles_vmec.py b/examples/particle_tracing/trace_particles_vmec.py index 14fb9f2..cbc3f34 100644 --- a/examples/particle_tracing/trace_particles_vmec.py +++ b/examples/particle_tracing/trace_particles_vmec.py @@ -12,7 +12,7 @@ # Input parameters tmax = 1e-4 timestep = 1.e-8 -times_to_trace=5000 +times_to_trace=1000 nparticles_per_core=6 nparticles = number_of_processors_to_use*nparticles_per_core n_particles_to_plot = 4 diff --git a/examples/particle_tracing/trace_particles_vmec_Electric_field.py b/examples/particle_tracing/trace_particles_vmec_Electric_field.py index 6aaa005..ce9e380 100644 --- a/examples/particle_tracing/trace_particles_vmec_Electric_field.py +++ b/examples/particle_tracing/trace_particles_vmec_Electric_field.py @@ -24,11 +24,11 @@ energy=FUSION_ALPHA_PARTICLE_ENERGY # Load coils and field -wout_file = os.path.join(os.path.dirname(__file__), "../input_files", "wout_LandremanPaul2021_QA_reactorScale_lowres.nc") +wout_file = os.path.join(os.path.dirname(__file__), "..", "input_files", "wout_LandremanPaul2021_QA_reactorScale_lowres.nc") vmec = Vmec(wout_file) #Load electric field -Er_file=os.path.join(os.path.dirname(__file__), '../input_files','Er.h5') +Er_file = os.path.join(os.path.dirname(__file__), "..", "input_files", "Er.h5") Electric_field=Electric_field_flux(Er_filename=Er_file,vmec=vmec) # Initialize particles diff --git a/examples/particle_tracing/trace_particles_vmec_classifier.py b/examples/particle_tracing/trace_particles_vmec_classifier.py new file mode 100644 index 0000000..10c2e9b --- /dev/null +++ b/examples/particle_tracing/trace_particles_vmec_classifier.py @@ -0,0 +1,82 @@ +import os +number_of_processors_to_use = 1 # Parallelization, this should divide nparticles +os.environ["XLA_FLAGS"] = f'--xla_force_host_platform_device_count={number_of_processors_to_use}' +from time import time +import jax.numpy as jnp +import matplotlib.pyplot as plt +from essos.fields import Vmec +from essos.constants import ALPHA_PARTICLE_MASS, ALPHA_PARTICLE_CHARGE, FUSION_ALPHA_PARTICLE_ENERGY +from essos.dynamics import Tracing, Particles +from essos.surfaces import SurfaceClassifier +import numpy as np + +# Input parameters +tmax = 1e-4 +timestep = 1.e-8 +times_to_trace=1000 +nparticles_per_core=6 +nparticles = number_of_processors_to_use*nparticles_per_core +n_particles_to_plot = 4 +s = 0.6 # s-coordinate: flux surface label +theta = jnp.linspace(0, 2*jnp.pi, nparticles) +phi = jnp.linspace(0, 2*jnp.pi/2/4, nparticles) +atol = 1e-8 +rtol = 1e-8 +energy=FUSION_ALPHA_PARTICLE_ENERGY + +# Load coils and field +wout_file = os.path.join(os.path.dirname(__file__), "../input_files", "wout_LandremanPaul2021_QA_reactorScale_lowres.nc") +vmec = Vmec(wout_file) +boundary=SurfaceClassifier(vmec.surface,h=0.1) + +# Initialize particles +Z0 = jnp.zeros(nparticles) +phi0 = jnp.zeros(nparticles) +initial_xyz=jnp.array([[s]*nparticles, theta, phi]).T +particles = Particles(initial_xyz=initial_xyz, mass=ALPHA_PARTICLE_MASS, + charge=ALPHA_PARTICLE_CHARGE, energy=energy, field=vmec) +# Trace in ESSOS +time0 = time() +tracing = Tracing(field=vmec, model='GuidingCenterAdaptative', particles=particles, maxtime=tmax, + timestep=timestep,times_to_trace=times_to_trace, atol=atol,rtol=rtol,boundary=boundary) +print(f"ESSOS tracing of {nparticles} particles during {tmax}s took {time()-time0:.2f} seconds") +print(f"Final loss fraction: {tracing.loss_fractions[-1]*100:.2f}%") +trajectories = tracing.trajectories + +# Plot trajectories, velocity parallel to the magnetic field, loss fractions and/or energy error +fig = plt.figure(figsize=(9, 8)) +ax1 = fig.add_subplot(221, projection='3d') +ax2 = fig.add_subplot(222) +ax3 = fig.add_subplot(223) +ax4 = fig.add_subplot(224) + +# Plot 5 random particles +## Plot trajectories in 3D +vmec.surface.plot(ax=ax1, show=False, alpha=0.4) +tracing.plot(ax=ax1, show=False, n_trajectories_plot=nparticles) +for i in np.random.choice(nparticles, size=n_particles_to_plot, replace=False): + trajectory = trajectories[i] + ## Plot energy error + ax2.plot(tracing.times[2:], jnp.abs(tracing.energy()[i][2:]-particles.energy)/particles.energy, label=f'Particle {i+1}') + ## Plot velocity parallel to the magnetic field + ax3.plot(tracing.times, trajectory[:, 3]/particles.total_speed, label=f'Particle {i+1}') + ## Plot s-coordinate + ax4.plot(tracing.times, trajectory[:,0], label=f'Particle {i+1}') + # ax4.set_ylabel(r'$s=\psi/\psi_b$') +## Plot loss fractions +#ax4.plot(tracing.times, tracing.loss_fractions) +#ax4.set_ylabel('Loss Fraction');ax4.set_ylim(0, 1);ax4.set_xscale('log') +ax2.set_xscale('log') +ax2.set_yscale('log') +ax2.set_ylabel('Relative Energy Error') +ax2.set_xlabel('Time (s)') +ax3.set_ylim(-1, 1) +ax3.set_ylabel(r'$v_{\parallel}/v$') +ax3.set_xlabel('Time (s)') +ax4.set_xlabel('Time (s)') +plt.tight_layout() +plt.show() + +# # Save results in vtk format to analyze in Paraview +# vmec.surface.to_vtk('surface') +# tracing.to_vtk('trajectories') diff --git a/examples/particle_tracing_collisions/statistics_collisions_velocity_distributions_mu_Adaptative.py b/examples/particle_tracing_collisions/statistics_collisions_velocity_distributions_mu_Adaptative.py index 33d67f0..70613ea 100644 --- a/examples/particle_tracing_collisions/statistics_collisions_velocity_distributions_mu_Adaptative.py +++ b/examples/particle_tracing_collisions/statistics_collisions_velocity_distributions_mu_Adaptative.py @@ -5,9 +5,8 @@ import jax.numpy as jnp import matplotlib.pyplot as plt import matplotlib.colors -from essos.fields import BiotSavart -from essos.coils import Coils -from essos.constants import PROTON_MASS, ONE_EV,ELECTRON_MASS,SPEED_OF_LIGHT +from essos.fields import BiotSavart,Vmec +from essos.constants import PROTON_MASS, ONE_EV,ELECTRON_MASS,SPEED_OF_LIGHT,ELEMENTARY_CHARGE from essos.dynamics import Tracing, Particles from essos.background_species import BackgroundSpecies,gamma_ab import numpy as np @@ -17,38 +16,42 @@ # to use higher precision config.update("jax_enable_x64", True) - - - # Input parameters -tmax = 1e-5 +tmax = 1e-4 dt=1.e-14 -times_to_trace=100 +times_to_trace=1000 nparticles_per_core=10 nparticles = number_of_processors_to_use*nparticles_per_core -R0 = 1.25#jnp.linspace(1.23, 1.27, nparticles) -atol = 1.e-6 -rtol=0. -rejected_steps=100 +s=0.25 +num_steps = jnp.round(tmax/dt) mass=PROTON_MASS mass_e=ELECTRON_MASS T_test=3000. energy=T_test*ONE_EV +# # Load coils and field +# json_file = os.path.join(os.path.dirname(__file__), '../input_files', 'ESSOS_biot_savart_LandremanPaulQA.json') +# coils = Coils_from_json(json_file) +plt.rcParams.update({'font.size': 16}) +# field = BiotSavart(coils) -# Load coils and field -json_file = os.path.join(os.path.dirname(__file__), '..', 'input_files', 'ESSOS_biot_savart_LandremanPaulQA.json') -coils = Coils.from_json(json_file) -field = BiotSavart(coils) +# # Initialize particles +# Z0 = jnp.zeros(nparticles) +# phi0 = jnp.zeros(nparticles) +# initial_xyz=jnp.array([R0*jnp.cos(phi0), R0*jnp.sin(phi0), Z0]).T +# particles = Particles(initial_xyz=initial_xyz,initial_vparallel_over_v=1.0*jnp.ones(nparticles), mass=mass, energy=energy) -# Initialize particles -Z0 = jnp.zeros(nparticles) -phi0 = jnp.zeros(nparticles) -initial_xyz=jnp.array([R0*jnp.cos(phi0), R0*jnp.sin(phi0), Z0]).T -particles = Particles(initial_xyz=initial_xyz,initial_vparallel_over_v=1.0*jnp.ones(nparticles), mass=mass, energy=energy) +# Load coils and field +wout_file = os.path.join(os.path.dirname(__file__), '..', 'input_files', "wout_LandremanPaul2021_QA_reactorScale_lowres.nc") +vmec = Vmec(wout_file, ntheta=60, nphi=60, range_torus='half period', close=True) +theta = jnp.zeros(nparticles) +phi = jnp.zeros(nparticles) +initial_xyz=jnp.array([[s]*nparticles, theta, phi]).T +particles = Particles(initial_xyz=initial_xyz, mass=mass, + charge=ELEMENTARY_CHARGE, energy=energy, field=vmec,initial_vparallel_over_v=1.0*jnp.ones(nparticles)) #Initialize background species number_species=1 #(electrons,deuterium) @@ -70,13 +73,13 @@ pitch_sigma=jnp.sqrt(2.**2/12) -# Trace in ESSOS time0 = time() -tracing = Tracing(field=field, model='GuidingCenterCollisionsMuAdaptative', particles=particles, - maxtime=tmax, timestep=dt,times_to_trace=times_to_trace, rtol=rtol,atol=atol,species=species,tag_gc=0.,rejected_steps=100) +tracing = Tracing(field=vmec, model='GuidingCenterCollisionsMuAdaptative', particles=particles, + maxtime=tmax, timestep=dt,times_to_trace=times_to_trace,species=species,tag_gc=0.) print(f"ESSOS tracing took {time()-time0:.2f} seconds") trajectories = tracing.trajectories + # Plot trajectories, velocity parallel to the magnetic field, and energy error fig = plt.figure(figsize=(9, 8)) ax1 = fig.add_subplot(221, projection='3d') @@ -84,88 +87,130 @@ ax3 = fig.add_subplot(223) ax4 = fig.add_subplot(224) -coils.plot(ax=ax1, show=False) -tracing.plot(ax=ax1, show=False) +#vmec.plot(ax=ax1, show=False) +#tracing.plot(ax=ax1, show=False) +# Plot only a random subset of 10 particles in 3D +subset_size = 10 +import numpy as np +subset_indices = np.random.choice(len(trajectories), subset_size, replace=False) -for i, trajectory in enumerate(trajectories): - ax2.plot(tracing.times, (tracing.energy[i]-tracing.energy[i,0])/tracing.energy[i,0], label=f'Particle {i+1}') - ax3.plot(tracing.times, trajectory[:, 3]*SPEED_OF_LIGHT/jnp.sqrt(tracing.energy[i]/mass*2.), label=f'Particle {i+1}') +for i in subset_indices: + trajectory = trajectories[i] + ax1.plot(trajectory[:,0], trajectory[:,1], trajectory[:,2], label=f'Particle {i+1}') + ax2.plot(tracing.times, (tracing.energy()[i]-tracing.energy()[i,0])/tracing.energy()[i,0], label=f'Particle {i+1}') + ax3.plot(tracing.times, 299792458*trajectory[:, 3]/jnp.sqrt(tracing.energy()[i]/mass*2.), label=f'Particle {i+1}') ax4.plot(jnp.sqrt(trajectory[:,0]**2+trajectory[:,1]**2), trajectory[:, 2], label=f'Particle {i+1}') - - -ax2.set_xlabel('Time (s)') -ax2.set_ylabel('Normalized energy variation') -ax3.set_ylabel(r'$v_{\parallel}/v$') -ax3.set_xlabel('Time (s)') -ax4.set_xlabel('R (m)') -ax4.set_ylabel('Z (m)') +# Set bold font for all axes and tick labels +for ax in [ax1, ax2, ax3, ax4]: + ax.xaxis.label.set_fontweight('bold') + ax.yaxis.label.set_fontweight('bold') + ax.title.set_fontweight('bold') + for label in (ax.get_xticklabels() + ax.get_yticklabels()): + label.set_fontweight('bold') + +ax2.set_xlabel(r'$t~[\mathrm{s}]$', fontweight='bold') +ax2.set_ylabel(r'$\frac{E-E_0}{E_0}$', fontweight='bold') +ax3.set_ylabel(r'$v_{\parallel}/v$', fontweight='bold') +ax3.set_xlabel(r'$t~[\mathrm{s}]$', fontweight='bold') +ax4.set_xlabel(r'$R~[\mathrm{m}]$', fontweight='bold') +ax4.set_ylabel(r'$Z~[\mathrm{m}]$', fontweight='bold') plt.tight_layout() plt.savefig('traj.pdf') - - -v=jnp.sqrt(tracing.energy*2./particles.mass) +v=jnp.sqrt(tracing.energy()*2./particles.mass) vpar=trajectories[:,:,3]*SPEED_OF_LIGHT -vperp=tracing.vperp_final +vpar=jnp.where(jnp.isfinite(vpar), vpar, jnp.nan) +vperp=tracing.v_perp() pitch=vpar/v -# Plot distribution in velocities initial t and final -fig3 = plt.figure(figsize=(9, 8)) -ax13 = fig3.add_subplot(241) -ax23 = fig3.add_subplot(242) -ax33 = fig3.add_subplot(243) -ax43 = fig3.add_subplot(244) -ax53 = fig3.add_subplot(245) -ax63 = fig3.add_subplot(246) -ax73 = fig3.add_subplot(247) -ax83 = fig3.add_subplot(248) -ax13.plot(tracing.times,jnp.nanmean(v/SPEED_OF_LIGHT,axis=0)) -ax13.axhline(y=v_mean, color='r', linestyle='--') -ax23.plot(tracing.times,jnp.nanstd(v/SPEED_OF_LIGHT,axis=0)) -ax23.axhline(y=v_sigma, color='r', linestyle='--') -ax33.plot(tracing.times,jnp.nanmean(pitch,axis=0)) -ax33.axhline(y=pitch_mean, color='r', linestyle='--') -ax43.plot(tracing.times,jnp.nanstd(pitch,axis=0)) -ax43.axhline(y=pitch_sigma, color='r', linestyle='--') -ax53.plot(tracing.times,jnp.nanmean(vpar/SPEED_OF_LIGHT,axis=0)) -ax53.axhline(y=vpar_mean, color='r', linestyle='--') -ax63.plot(tracing.times,jnp.nanstd(vpar/SPEED_OF_LIGHT,axis=0)) -ax63.axhline(y=vpar_sigma, color='r', linestyle='--') -ax73.plot(tracing.times,jnp.nanmean(vperp/SPEED_OF_LIGHT,axis=0)) -ax73.axhline(y=vperp_mean, color='r', linestyle='--') -ax83.plot(tracing.times,jnp.nanstd(vperp/SPEED_OF_LIGHT,axis=0)) -ax83.axhline(y=vperp_sigma, color='r', linestyle='--') -ax13.set_title('Mean energy') -ax13.set_xlabel('time') -ax13.set_ylabel('Energy') -ax23.set_title('sigma energy') -ax23.set_xlabel('time') -ax23.set_ylabel('Energy') -ax33.set_title('Mean pitch') -ax33.set_xlabel('time') -ax33.set_ylabel('pitch') -ax43.set_title('sigma pitch') -ax43.set_xlabel('time') -ax43.set_ylabel('pitch') -ax53.set_title('Mean vpar') -ax53.set_xlabel('time') -ax53.set_ylabel('vpar') -ax63.set_title('sigma vpar') -ax63.set_xlabel('time') -ax63.set_ylabel('vpar') -ax73.set_title('Mean vperp') -ax73.set_xlabel('time') -ax73.set_ylabel('vperp') -ax83.set_title('sigma vperp') -ax83.set_xlabel('time') -ax83.set_ylabel('vperp') -plt.tight_layout() -plt.savefig('statistics.pdf') +# Improve font size for all plots +plt.rcParams.update({'font.size': 18, 'font.weight': 'bold'}) + +# 1. v +fig_v = plt.figure(figsize=(7, 5)) +ax_v_mean = fig_v.add_subplot(211) +ax_v_std = fig_v.add_subplot(212) +for ax in [ax_v_mean, ax_v_std]: + for label in (ax.get_xticklabels() + ax.get_yticklabels()): + label.set_fontweight('bold') +ax_v_mean.plot(tracing.times, jnp.nanmean(v/SPEED_OF_LIGHT, axis=0), linewidth=5) +ax_v_mean.axhline(y=v_mean, color='r', linestyle='--', linewidth=5) +ax_v_mean.set_title(r'$\langle v \rangle$', fontweight='bold') +ax_v_mean.set_xlabel('time', fontweight='bold') +ax_v_mean.set_ylabel(r'$v/c$', fontweight='bold') +ax_v_std.plot(tracing.times, jnp.nanstd(v/SPEED_OF_LIGHT, axis=0), linewidth=5) +ax_v_std.axhline(y=v_sigma, color='r', linestyle='--', linewidth=5) +ax_v_std.set_title(r'$\sigma(v)$', fontweight='bold') +ax_v_std.set_xlabel('time', fontweight='bold') +ax_v_std.set_ylabel(r'$v/c$', fontweight='bold') +plt.tight_layout() +fig_v.savefig('statistics_v.pdf', dpi=300) + +# 2. pitch +fig_pitch = plt.figure(figsize=(7, 5)) +ax_pitch_mean = fig_pitch.add_subplot(211) +ax_pitch_std = fig_pitch.add_subplot(212) +for ax in [ax_pitch_mean, ax_pitch_std]: + for label in (ax.get_xticklabels() + ax.get_yticklabels()): + label.set_fontweight('bold') +ax_pitch_mean.plot(tracing.times, jnp.nanmean(pitch, axis=0), linewidth=5) +ax_pitch_mean.axhline(y=pitch_mean, color='r', linestyle='--', linewidth=5) +ax_pitch_mean.set_title(r'$\langle \text{pitch} \rangle$', fontweight='bold') +ax_pitch_mean.set_xlabel('time', fontweight='bold') +ax_pitch_mean.set_ylabel('pitch', fontweight='bold') +ax_pitch_std.plot(tracing.times, jnp.nanstd(pitch, axis=0), linewidth=5) +ax_pitch_std.axhline(y=pitch_sigma, color='r', linestyle='--', linewidth=5) +ax_pitch_std.set_title(r'$\sigma(\text{pitch})$', fontweight='bold') +ax_pitch_std.set_xlabel('time', fontweight='bold') +ax_pitch_std.set_ylabel('pitch', fontweight='bold') +plt.tight_layout() +fig_pitch.savefig('statistics_pitch.pdf', dpi=300) + +# 3. v_parallel/c +fig_vpar = plt.figure(figsize=(7, 5)) +ax_vpar_mean = fig_vpar.add_subplot(211) +ax_vpar_std = fig_vpar.add_subplot(212) +for ax in [ax_vpar_mean, ax_vpar_std]: + for label in (ax.get_xticklabels() + ax.get_yticklabels()): + label.set_fontweight('bold') +ax_vpar_mean.plot(tracing.times, jnp.nanmean(vpar/SPEED_OF_LIGHT, axis=0), linewidth=5) +ax_vpar_mean.axhline(y=vpar_mean, color='r', linestyle='--', linewidth=5) +ax_vpar_mean.set_title(r'$\langle v_{\parallel}/c \rangle$', fontweight='bold') +ax_vpar_mean.set_xlabel(r'$t~[\mathrm{s}]$', fontweight='bold') +ax_vpar_mean.set_ylabel(r'$v_{\parallel}/c$', fontweight='bold') +ax_vpar_std.plot(tracing.times, jnp.nanstd(vpar/SPEED_OF_LIGHT, axis=0), linewidth=5) +ax_vpar_std.axhline(y=vpar_sigma, color='r', linestyle='--', linewidth=5) +ax_vpar_std.set_title(r'$\sigma(v_{\parallel}/c)$', fontweight='bold') +ax_vpar_std.set_xlabel(r'$t~[\mathrm{s}]$', fontweight='bold') +ax_vpar_std.set_ylabel(r'$\sigma_{v_{\parallel}/c}$', fontweight='bold') +plt.tight_layout() +fig_vpar.savefig('statistics_vpar.pdf', dpi=300) + +# 4. v_perp/c +fig_vperp = plt.figure(figsize=(7, 5)) +ax_vperp_mean = fig_vperp.add_subplot(211) +ax_vperp_std = fig_vperp.add_subplot(212) +for ax in [ax_vperp_mean, ax_vperp_std]: + for label in (ax.get_xticklabels() + ax.get_yticklabels()): + label.set_fontweight('bold') +ax_vperp_mean.plot(tracing.times, jnp.nanmean(vperp/SPEED_OF_LIGHT, axis=0), linewidth=5) +ax_vperp_mean.axhline(y=vperp_mean, color='r', linestyle='--', linewidth=5) +ax_vperp_mean.set_title(r'$\langle v_{\perp}/c \rangle$', fontweight='bold') +ax_vperp_mean.set_xlabel(r'$t~[\mathrm{s}]$', fontweight='bold') +ax_vperp_mean.set_ylabel(r'$v_{\perp}/c$', fontweight='bold') +ax_vperp_std.plot(tracing.times, jnp.nanstd(vperp/SPEED_OF_LIGHT, axis=0), linewidth=5) +ax_vperp_std.axhline(y=vperp_sigma, color='r', linestyle='--', linewidth=5) +ax_vperp_std.set_title(r'$\sigma(v_{\perp}/c)$', fontweight='bold') +ax_vperp_std.set_xlabel(r'$t~[\mathrm{s}]$', fontweight='bold') +ax_vperp_std.set_ylabel(r'$\sigma_{v_{\perp}/c}$', fontweight='bold') +plt.tight_layout() +fig_vperp.savefig('statistics_vperp.pdf', dpi=300) + # Plot distribution in velocities initial t and final fig2 = plt.figure(figsize=(9, 8)) @@ -179,19 +224,16 @@ ax82 = fig2.add_subplot(258) nbins=64 -v0=jnp.sqrt(tracing.energy[:,0]*2./particles.mass) -vfinal=jnp.sqrt(tracing.energy[:,-1]*2./particles.mass) -vperp0=tracing.vperp_final[:,0] -vperpfinal=tracing.vperp_final[:,-1] -vpar0=trajectories[:,0,3] -vparfinal=trajectories[:,-1,3] +v0=jnp.sqrt(tracing.energy()[:,0]*2./particles.mass)/SPEED_OF_LIGHT +vfinal=jnp.sqrt(tracing.energy()[:,-1]*2./particles.mass)/SPEED_OF_LIGHT +vperp0=tracing.v_perp()[:,0]/SPEED_OF_LIGHT +vperpfinal=tracing.v_perp()[:,-1]/SPEED_OF_LIGHT +vpar0=vpar[:,0]/SPEED_OF_LIGHT +vparfinal=vpar[:,-1]/SPEED_OF_LIGHT pitch0=vpar0/v0 pitch_final=vparfinal/vfinal - - - bad_indices_v0 = jnp.isnan(v0) bad_indices_vfinal = jnp.isnan(vfinal) bad_indices_pitch0 = jnp.isnan(pitch0) @@ -242,33 +284,53 @@ ax62.stairs(pitch_tfinal_counts,pitch_tfinal_bins) ax72.stairs(vperp_t0_counts,vperp_t0_bins) ax82.stairs(vperp_tfinal_counts,vperp_tfinal_bins) - -ax12.set_title('t=0') -ax12.set_xlabel('v') -ax12.set_ylabel('Counts') -ax22.set_title('t=t_final') -ax22.set_xlabel('v') -ax22.set_ylabel('Counts') -ax32.set_title('t=0') -ax32.set_ylabel('Counts') -ax32.set_xlabel(r'$v_{\parallel}$') -ax42.set_title('t=t_final') -ax42.set_xlabel(r'$v_{parallel}$') -ax42.set_ylabel('Counts') -ax52.set_title('t=0') -ax52.set_xlabel(r'$v_{\parallel}/v$') -ax52.set_ylabel('Counts') -ax62.set_title('t=t_final') -ax62.set_xlabel(r'$v_{\parallel}/v$') -ax62.set_ylabel('Counts') -ax72.set_title('t=0') -ax72.set_ylabel('Counts') -ax72.set_xlabel(r'$v_{\perp}$') -ax82.set_title('t=t_final') -ax82.set_ylabel('Counts') -ax82.set_xlabel(r'$v_{\perp}$') +plt.figure(figsize=(7, 5)) +plt.hist(good_vfinal, bins=nbins, color='b', edgecolor='black', alpha=0.7) +plt.axvline(np.mean(good_v0), color='r', linestyle='--', linewidth=3, label='Initial Mean') +plt.title(r'$v/c$ Distribution', fontweight='bold') +plt.xlabel(r'$v/c$', fontweight='bold') +plt.ylabel('Counts', fontweight='bold') +plt.legend(fontsize=14) +plt.tight_layout() +plt.savefig('dist_v.pdf', dpi=300) + +plt.figure(figsize=(7, 5)) +plt.hist(good_pitch_final, bins=nbins, color='g', edgecolor='black', alpha=0.7) +plt.axvline(np.mean(good_pitch0), color='r', linestyle='--', linewidth=3, label='Initial Mean') +plt.title(r'Pitch Distribution', fontweight='bold') +plt.xlabel(r'Pitch', fontweight='bold') +plt.ylabel('Counts', fontweight='bold') +plt.legend(fontsize=14) +plt.tight_layout() +plt.savefig('dist_pitch.pdf', dpi=300) + +plt.figure(figsize=(7, 5)) +plt.hist(good_vpar_final, bins=nbins, color='#FA7000', edgecolor='black', alpha=0.7) +plt.axvline(np.mean(good_vpar0), color='b', linestyle='--', linewidth=3, label='Initial Mean') +plt.title(r'$v_{\parallel}/c$ Distribution', fontweight='bold') +plt.xlabel(r'$v_{\parallel}/c$', fontweight='bold') +plt.ylabel('Counts', fontweight='bold') +plt.legend(fontsize=14) +plt.tight_layout() +plt.savefig('dist_vpar.pdf', dpi=300) + +plt.figure(figsize=(7, 5)) +plt.hist(good_vperp_final, bins=nbins, color='m', edgecolor='black', alpha=0.7) +plt.axvline(np.mean(good_vperp0), color='b', linestyle='--', linewidth=3, label='Initial Mean') +plt.title(r'$v_{\perp}/c$ Distribution', fontweight='bold') +plt.xlabel(r'$v_{\perp}/c$', fontweight='bold') +plt.ylabel('Counts', fontweight='bold') +plt.legend(fontsize=14) +plt.tight_layout() +plt.savefig('dist_vperp.pdf', dpi=300) +plt.figure(figsize=(7, 5)) +plt.hist(good_vperp_final, bins=nbins, color='#FA7000', edgecolor='black', alpha=0.7) +plt.axvline(np.mean(good_vperp0), color='b', linestyle='--', linewidth=3, label='Initial Mean') +plt.title(r'$v_{\perp}/c$ Distribution', fontweight='bold') +plt.xlabel(r'$v_{\perp}/c$', fontweight='bold') +plt.ylabel('Counts', fontweight='bold') +plt.legend(fontsize=14) plt.tight_layout() -plt.savefig('dist.pdf') - +plt.savefig('dist_vperp_color.pdf', dpi=300) diff --git a/examples/particle_tracing_collisions/statistics_collisions_velocity_distributions_mu_Fixed.py b/examples/particle_tracing_collisions/statistics_collisions_velocity_distributions_mu_Fixed.py index aa1a6d3..41d3e03 100644 --- a/examples/particle_tracing_collisions/statistics_collisions_velocity_distributions_mu_Fixed.py +++ b/examples/particle_tracing_collisions/statistics_collisions_velocity_distributions_mu_Fixed.py @@ -5,9 +5,8 @@ import jax.numpy as jnp import matplotlib.pyplot as plt import matplotlib.colors -from essos.fields import BiotSavart -from essos.coils import Coils -from essos.constants import PROTON_MASS, ONE_EV,ELECTRON_MASS,SPEED_OF_LIGHT +from essos.fields import BiotSavart,Vmec +from essos.constants import PROTON_MASS, ONE_EV,ELECTRON_MASS,SPEED_OF_LIGHT,ELEMENTARY_CHARGE from essos.dynamics import Tracing, Particles from essos.background_species import BackgroundSpecies,gamma_ab import numpy as np @@ -18,12 +17,12 @@ config.update("jax_enable_x64", True) # Input parameters -tmax = 1.e-5 +tmax = 1.e-4 dt=1.e-8 -times_to_trace=100 +times_to_trace=1000 nparticles_per_core=10 nparticles = number_of_processors_to_use*nparticles_per_core -R0 = 1.25 +s=0.25 num_steps = jnp.round(tmax/dt) mass=PROTON_MASS mass_e=ELECTRON_MASS @@ -31,17 +30,29 @@ energy=T_test*ONE_EV +# # Load coils and field +# json_file = os.path.join(os.path.dirname(__file__), '../input_files', 'ESSOS_biot_savart_LandremanPaulQA.json') +# coils = Coils_from_json(json_file) +plt.rcParams.update({'font.size': 16}) +# field = BiotSavart(coils) + +# # Initialize particles +# Z0 = jnp.zeros(nparticles) +# phi0 = jnp.zeros(nparticles) +# initial_xyz=jnp.array([R0*jnp.cos(phi0), R0*jnp.sin(phi0), Z0]).T +# particles = Particles(initial_xyz=initial_xyz,initial_vparallel_over_v=1.0*jnp.ones(nparticles), mass=mass, energy=energy) + + # Load coils and field -json_file = os.path.join(os.path.dirname(__file__), '..', 'input_files', 'ESSOS_biot_savart_LandremanPaulQA.json') -coils = Coils.from_json(json_file) -field = BiotSavart(coils) +wout_file = os.path.join(os.path.dirname(__file__), '..', 'input_files', "wout_LandremanPaul2021_QA_reactorScale_lowres.nc") +vmec = Vmec(wout_file, ntheta=60, nphi=60, range_torus='half period', close=True) -# Initialize particles -Z0 = jnp.zeros(nparticles) -phi0 = jnp.zeros(nparticles) -initial_xyz=jnp.array([R0*jnp.cos(phi0), R0*jnp.sin(phi0), Z0]).T -particles = Particles(initial_xyz=initial_xyz,initial_vparallel_over_v=1.0*jnp.ones(nparticles), mass=mass, energy=energy) +theta = jnp.zeros(nparticles) +phi = jnp.zeros(nparticles) +initial_xyz=jnp.array([[s]*nparticles, theta, phi]).T +particles = Particles(initial_xyz=initial_xyz, mass=mass, + charge=ELEMENTARY_CHARGE, energy=energy, field=vmec,initial_vparallel_over_v=1.0*jnp.ones(nparticles)) #Initialize background species number_species=1 #(electrons,deuterium) @@ -63,13 +74,13 @@ pitch_sigma=jnp.sqrt(2.**2/12) -# Trace in ESSOS time0 = time() -tracing = Tracing(field=field, model='GuidingCenterCollisionsMuFixed', particles=particles, +tracing = Tracing(field=vmec, model='GuidingCenterCollisionsMuFixed', particles=particles, maxtime=tmax, timestep=dt,times_to_trace=times_to_trace,species=species,tag_gc=0.) print(f"ESSOS tracing took {time()-time0:.2f} seconds") trajectories = tracing.trajectories + # Plot trajectories, velocity parallel to the magnetic field, and energy error fig = plt.figure(figsize=(9, 8)) ax1 = fig.add_subplot(221, projection='3d') @@ -77,86 +88,128 @@ ax3 = fig.add_subplot(223) ax4 = fig.add_subplot(224) -coils.plot(ax=ax1, show=False) -tracing.plot(ax=ax1, show=False) - -for i, trajectory in enumerate(trajectories): - ax2.plot(tracing.times, (tracing.energy[i]-tracing.energy[i,0])/tracing.energy[i,0], label=f'Particle {i+1}') - ax3.plot(tracing.times, 299792458*trajectory[:, 3]/jnp.sqrt(tracing.energy[i]/mass*2.), label=f'Particle {i+1}') - ax4.plot(jnp.sqrt(trajectory[:,0]**2+trajectory[:,1]**2), trajectory[:, 2], label=f'Particle {i+1}') +#vmec.plot(ax=ax1, show=False) +#tracing.plot(ax=ax1, show=False) +# Plot only a random subset of 10 particles in 3D +subset_size = 10 +import numpy as np +subset_indices = np.random.choice(len(trajectories), subset_size, replace=False) +for i in subset_indices: + trajectory = trajectories[i] + ax1.plot(trajectory[:,0], trajectory[:,1], trajectory[:,2], label=f'Particle {i+1}') + ax2.plot(tracing.times, (tracing.energy()[i]-tracing.energy()[i,0])/tracing.energy()[i,0], label=f'Particle {i+1}') + ax3.plot(tracing.times, 299792458*trajectory[:, 3]/jnp.sqrt(tracing.energy()[i]/mass*2.), label=f'Particle {i+1}') + ax4.plot(jnp.sqrt(trajectory[:,0]**2+trajectory[:,1]**2), trajectory[:, 2], label=f'Particle {i+1}') -ax2.set_xlabel('Time (s)') -ax2.set_ylabel('Normalized energy variation') -ax3.set_ylabel(r'$v_{\parallel}/v$') -ax3.set_xlabel('Time (s)') -ax4.set_xlabel('R (m)') -ax4.set_ylabel('Z (m)') +# Set bold font for all axes and tick labels +for ax in [ax1, ax2, ax3, ax4]: + ax.xaxis.label.set_fontweight('bold') + ax.yaxis.label.set_fontweight('bold') + ax.title.set_fontweight('bold') + for label in (ax.get_xticklabels() + ax.get_yticklabels()): + label.set_fontweight('bold') + +ax2.set_xlabel(r'$t~[\mathrm{s}]$', fontweight='bold') +ax2.set_ylabel(r'$\frac{E-E_0}{E_0}$', fontweight='bold') +ax3.set_ylabel(r'$v_{\parallel}/v$', fontweight='bold') +ax3.set_xlabel(r'$t~[\mathrm{s}]$', fontweight='bold') +ax4.set_xlabel(r'$R~[\mathrm{m}]$', fontweight='bold') +ax4.set_ylabel(r'$Z~[\mathrm{m}]$', fontweight='bold') plt.tight_layout() plt.savefig('traj.pdf') - - - -v=jnp.sqrt(tracing.energy*2./particles.mass) -vpar=trajectories[:,:,3] -vperp=tracing.vperp_final +v=jnp.sqrt(tracing.energy()*2./particles.mass) +vpar=trajectories[:,:,3]*SPEED_OF_LIGHT +vpar=jnp.where(jnp.isfinite(vpar), vpar, jnp.nan) +vperp=tracing.v_perp() pitch=vpar/v -# Plot distribution in velocities initial t and final -fig3 = plt.figure(figsize=(9, 8)) -ax13 = fig3.add_subplot(241) -ax23 = fig3.add_subplot(242) -ax33 = fig3.add_subplot(243) -ax43 = fig3.add_subplot(244) -ax53 = fig3.add_subplot(245) -ax63 = fig3.add_subplot(246) -ax73 = fig3.add_subplot(247) -ax83 = fig3.add_subplot(248) -ax13.plot(tracing.times,jnp.nanmean(v/SPEED_OF_LIGHT,axis=0)) -ax13.axhline(y=v_mean, color='r', linestyle='--') -ax23.plot(tracing.times,jnp.nanstd(v/SPEED_OF_LIGHT,axis=0)) -ax23.axhline(y=v_sigma, color='r', linestyle='--') -ax33.plot(tracing.times,jnp.nanmean(pitch,axis=0)) -ax33.axhline(y=pitch_mean, color='r', linestyle='--') -ax43.plot(tracing.times,jnp.nanstd(pitch,axis=0)) -ax43.axhline(y=pitch_sigma, color='r', linestyle='--') -ax53.plot(tracing.times,jnp.nanmean(vpar/SPEED_OF_LIGHT,axis=0)) -ax53.axhline(y=vpar_mean, color='r', linestyle='--') -ax63.plot(tracing.times,jnp.nanstd(vpar/SPEED_OF_LIGHT,axis=0)) -ax63.axhline(y=vpar_sigma, color='r', linestyle='--') -ax73.plot(tracing.times,jnp.nanmean(vperp/SPEED_OF_LIGHT,axis=0)) -ax73.axhline(y=vperp_mean, color='r', linestyle='--') -ax83.plot(tracing.times,jnp.nanstd(vperp/SPEED_OF_LIGHT,axis=0)) -ax83.axhline(y=vperp_sigma, color='r', linestyle='--') -ax13.set_title('Mean energy') -ax13.set_xlabel('time') -ax13.set_ylabel('Energy') -ax23.set_title('sigma energy') -ax23.set_xlabel('time') -ax23.set_ylabel('Energy') -ax33.set_title('Mean pitch') -ax33.set_xlabel('time') -ax33.set_ylabel('pitch') -ax43.set_title('sigma pitch') -ax43.set_xlabel('time') -ax43.set_ylabel('pitch') -ax53.set_title('Mean vpar') -ax53.set_xlabel('time') -ax53.set_ylabel('vpar') -ax63.set_title('sigma vpar') -ax63.set_xlabel('time') -ax63.set_ylabel('vpar') -ax73.set_title('Mean vperp') -ax73.set_xlabel('time') -ax73.set_ylabel('vperp') -ax83.set_title('sigma vperp') -ax83.set_xlabel('time') -ax83.set_ylabel('vperp') + +# Improve font size for all plots +plt.rcParams.update({'font.size': 18, 'font.weight': 'bold'}) + +# 1. v +fig_v = plt.figure(figsize=(7, 5)) +ax_v_mean = fig_v.add_subplot(211) +ax_v_std = fig_v.add_subplot(212) +for ax in [ax_v_mean, ax_v_std]: + for label in (ax.get_xticklabels() + ax.get_yticklabels()): + label.set_fontweight('bold') +ax_v_mean.plot(tracing.times, jnp.nanmean(v/SPEED_OF_LIGHT, axis=0), linewidth=5) +ax_v_mean.axhline(y=v_mean, color='r', linestyle='--', linewidth=5) +ax_v_mean.set_title(r'$\langle v \rangle$', fontweight='bold') +ax_v_mean.set_xlabel('time', fontweight='bold') +ax_v_mean.set_ylabel(r'$v/c$', fontweight='bold') +ax_v_std.plot(tracing.times, jnp.nanstd(v/SPEED_OF_LIGHT, axis=0), linewidth=5) +ax_v_std.axhline(y=v_sigma, color='r', linestyle='--', linewidth=5) +ax_v_std.set_title(r'$\sigma(v)$', fontweight='bold') +ax_v_std.set_xlabel('time', fontweight='bold') +ax_v_std.set_ylabel(r'$v/c$', fontweight='bold') +plt.tight_layout() +fig_v.savefig('statistics_v.pdf', dpi=300) + +# 2. pitch +fig_pitch = plt.figure(figsize=(7, 5)) +ax_pitch_mean = fig_pitch.add_subplot(211) +ax_pitch_std = fig_pitch.add_subplot(212) +for ax in [ax_pitch_mean, ax_pitch_std]: + for label in (ax.get_xticklabels() + ax.get_yticklabels()): + label.set_fontweight('bold') +ax_pitch_mean.plot(tracing.times, jnp.nanmean(pitch, axis=0), linewidth=5) +ax_pitch_mean.axhline(y=pitch_mean, color='r', linestyle='--', linewidth=5) +ax_pitch_mean.set_title(r'$\langle \text{pitch} \rangle$', fontweight='bold') +ax_pitch_mean.set_xlabel('time', fontweight='bold') +ax_pitch_mean.set_ylabel('pitch', fontweight='bold') +ax_pitch_std.plot(tracing.times, jnp.nanstd(pitch, axis=0), linewidth=5) +ax_pitch_std.axhline(y=pitch_sigma, color='r', linestyle='--', linewidth=5) +ax_pitch_std.set_title(r'$\sigma(\text{pitch})$', fontweight='bold') +ax_pitch_std.set_xlabel('time', fontweight='bold') +ax_pitch_std.set_ylabel('pitch', fontweight='bold') +plt.tight_layout() +fig_pitch.savefig('statistics_pitch.pdf', dpi=300) + +# 3. v_parallel/c +fig_vpar = plt.figure(figsize=(7, 5)) +ax_vpar_mean = fig_vpar.add_subplot(211) +ax_vpar_std = fig_vpar.add_subplot(212) +for ax in [ax_vpar_mean, ax_vpar_std]: + for label in (ax.get_xticklabels() + ax.get_yticklabels()): + label.set_fontweight('bold') +ax_vpar_mean.plot(tracing.times, jnp.nanmean(vpar/SPEED_OF_LIGHT, axis=0), linewidth=5) +ax_vpar_mean.axhline(y=vpar_mean, color='r', linestyle='--', linewidth=5) +ax_vpar_mean.set_title(r'$\langle v_{\parallel}/c \rangle$', fontweight='bold') +ax_vpar_mean.set_xlabel(r'$t~[\mathrm{s}]$', fontweight='bold') +ax_vpar_mean.set_ylabel(r'$v_{\parallel}/c$', fontweight='bold') +ax_vpar_std.plot(tracing.times, jnp.nanstd(vpar/SPEED_OF_LIGHT, axis=0), linewidth=5) +ax_vpar_std.axhline(y=vpar_sigma, color='r', linestyle='--', linewidth=5) +ax_vpar_std.set_title(r'$\sigma(v_{\parallel}/c)$', fontweight='bold') +ax_vpar_std.set_xlabel(r'$t~[\mathrm{s}]$', fontweight='bold') +ax_vpar_std.set_ylabel(r'$\sigma_{v_{\parallel}/c}$', fontweight='bold') plt.tight_layout() -plt.savefig('statistics.pdf') +fig_vpar.savefig('statistics_vpar.pdf', dpi=300) + +# 4. v_perp/c +fig_vperp = plt.figure(figsize=(7, 5)) +ax_vperp_mean = fig_vperp.add_subplot(211) +ax_vperp_std = fig_vperp.add_subplot(212) +for ax in [ax_vperp_mean, ax_vperp_std]: + for label in (ax.get_xticklabels() + ax.get_yticklabels()): + label.set_fontweight('bold') +ax_vperp_mean.plot(tracing.times, jnp.nanmean(vperp/SPEED_OF_LIGHT, axis=0), linewidth=5) +ax_vperp_mean.axhline(y=vperp_mean, color='r', linestyle='--', linewidth=5) +ax_vperp_mean.set_title(r'$\langle v_{\perp}/c \rangle$', fontweight='bold') +ax_vperp_mean.set_xlabel(r'$t~[\mathrm{s}]$', fontweight='bold') +ax_vperp_mean.set_ylabel(r'$v_{\perp}/c$', fontweight='bold') +ax_vperp_std.plot(tracing.times, jnp.nanstd(vperp/SPEED_OF_LIGHT, axis=0), linewidth=5) +ax_vperp_std.axhline(y=vperp_sigma, color='r', linestyle='--', linewidth=5) +ax_vperp_std.set_title(r'$\sigma(v_{\perp}/c)$', fontweight='bold') +ax_vperp_std.set_xlabel(r'$t~[\mathrm{s}]$', fontweight='bold') +ax_vperp_std.set_ylabel(r'$\sigma_{v_{\perp}/c}$', fontweight='bold') +plt.tight_layout() +fig_vperp.savefig('statistics_vperp.pdf', dpi=300) @@ -172,19 +225,16 @@ ax82 = fig2.add_subplot(258) nbins=64 -v0=jnp.sqrt(tracing.energy[:,0]*2./particles.mass) -vfinal=jnp.sqrt(tracing.energy[:,-1]*2./particles.mass) -vperp0=tracing.vperp_final[:,0] -vperpfinal=tracing.vperp_final[:,-1] -vpar0=trajectories[:,0,3] -vparfinal=trajectories[:,-1,3] +v0=jnp.sqrt(tracing.energy()[:,0]*2./particles.mass)/SPEED_OF_LIGHT +vfinal=jnp.sqrt(tracing.energy()[:,-1]*2./particles.mass)/SPEED_OF_LIGHT +vperp0=tracing.v_perp()[:,0]/SPEED_OF_LIGHT +vperpfinal=tracing.v_perp()[:,-1]/SPEED_OF_LIGHT +vpar0=vpar[:,0]/SPEED_OF_LIGHT +vparfinal=vpar[:,-1]/SPEED_OF_LIGHT pitch0=vpar0/v0 pitch_final=vparfinal/vfinal - - - bad_indices_v0 = jnp.isnan(v0) bad_indices_vfinal = jnp.isnan(vfinal) bad_indices_pitch0 = jnp.isnan(pitch0) @@ -235,32 +285,53 @@ ax62.stairs(pitch_tfinal_counts,pitch_tfinal_bins) ax72.stairs(vperp_t0_counts,vperp_t0_bins) ax82.stairs(vperp_tfinal_counts,vperp_tfinal_bins) - -ax12.set_title('t=0') -ax12.set_xlabel('v') -ax12.set_ylabel('Counts') -ax22.set_title('t=t_final') -ax22.set_xlabel('v') -ax22.set_ylabel('Counts') -ax32.set_title('t=0') -ax32.set_ylabel('Counts') -ax32.set_xlabel(r'$v_{\parallel}$') -ax42.set_title('t=t_final') -ax42.set_xlabel(r'$v_{parallel}$') -ax42.set_ylabel('Counts') -ax52.set_title('t=0') -ax52.set_xlabel(r'$v_{\parallel}/v$') -ax52.set_ylabel('Counts') -ax62.set_title('t=t_final') -ax62.set_xlabel(r'$v_{\parallel}/v$') -ax62.set_ylabel('Counts') -ax72.set_title('t=0') -ax72.set_ylabel('Counts') -ax72.set_xlabel(r'$v_{\perp}$') -ax82.set_title('t=t_final') -ax82.set_ylabel('Counts') -ax82.set_xlabel(r'$v_{\perp}$') +plt.figure(figsize=(7, 5)) +plt.hist(good_vfinal, bins=nbins, color='b', edgecolor='black', alpha=0.7) +plt.axvline(np.mean(good_v0), color='r', linestyle='--', linewidth=3, label='Initial Mean') +plt.title(r'$v/c$ Distribution', fontweight='bold') +plt.xlabel(r'$v/c$', fontweight='bold') +plt.ylabel('Counts', fontweight='bold') +plt.legend(fontsize=14) +plt.tight_layout() +plt.savefig('dist_v.pdf', dpi=300) + +plt.figure(figsize=(7, 5)) +plt.hist(good_pitch_final, bins=nbins, color='g', edgecolor='black', alpha=0.7) +plt.axvline(np.mean(good_pitch0), color='r', linestyle='--', linewidth=3, label='Initial Mean') +plt.title(r'Pitch Distribution', fontweight='bold') +plt.xlabel(r'Pitch', fontweight='bold') +plt.ylabel('Counts', fontweight='bold') +plt.legend(fontsize=14) +plt.tight_layout() +plt.savefig('dist_pitch.pdf', dpi=300) + +plt.figure(figsize=(7, 5)) +plt.hist(good_vpar_final, bins=nbins, color='#FA7000', edgecolor='black', alpha=0.7) +plt.axvline(np.mean(good_vpar0), color='b', linestyle='--', linewidth=3, label='Initial Mean') +plt.title(r'$v_{\parallel}/c$ Distribution', fontweight='bold') +plt.xlabel(r'$v_{\parallel}/c$', fontweight='bold') +plt.ylabel('Counts', fontweight='bold') +plt.legend(fontsize=14) +plt.tight_layout() +plt.savefig('dist_vpar.pdf', dpi=300) + +plt.figure(figsize=(7, 5)) +plt.hist(good_vperp_final, bins=nbins, color='m', edgecolor='black', alpha=0.7) +plt.axvline(np.mean(good_vperp0), color='b', linestyle='--', linewidth=3, label='Initial Mean') +plt.title(r'$v_{\perp}/c$ Distribution', fontweight='bold') +plt.xlabel(r'$v_{\perp}/c$', fontweight='bold') +plt.ylabel('Counts', fontweight='bold') +plt.legend(fontsize=14) +plt.tight_layout() +plt.savefig('dist_vperp.pdf', dpi=300) +plt.figure(figsize=(7, 5)) +plt.hist(good_vperp_final, bins=nbins, color='#FA7000', edgecolor='black', alpha=0.7) +plt.axvline(np.mean(good_vperp0), color='b', linestyle='--', linewidth=3, label='Initial Mean') +plt.title(r'$v_{\perp}/c$ Distribution', fontweight='bold') +plt.xlabel(r'$v_{\perp}/c$', fontweight='bold') +plt.ylabel('Counts', fontweight='bold') +plt.legend(fontsize=14) plt.tight_layout() -plt.savefig('dist.pdf') +plt.savefig('dist_vperp_color.pdf', dpi=300) diff --git a/examples/particle_tracing_collisions/statistics_collisions_velocity_distributions_mu_time.py b/examples/particle_tracing_collisions/statistics_collisions_velocity_distributions_mu_time.py index bab9b8b..ec10d27 100644 --- a/examples/particle_tracing_collisions/statistics_collisions_velocity_distributions_mu_time.py +++ b/examples/particle_tracing_collisions/statistics_collisions_velocity_distributions_mu_time.py @@ -5,9 +5,8 @@ import jax.numpy as jnp import matplotlib.pyplot as plt import matplotlib.colors -from essos.fields import BiotSavart -from essos.coils import Coils -from essos.constants import PROTON_MASS, ONE_EV,ELECTRON_MASS,SPEED_OF_LIGHT +from essos.fields import BiotSavart,Vmec +from essos.constants import PROTON_MASS, ONE_EV,ELECTRON_MASS,SPEED_OF_LIGHT,ELEMENTARY_CHARGE from essos.dynamics import Tracing, Particles from essos.background_species import BackgroundSpecies,gamma_ab import numpy as np @@ -15,13 +14,13 @@ # Input parameters light_speed=SPEED_OF_LIGHT -tmax = 1.e-5 +tmax = 1.e-4 dt=1.e-8 -nparticles_per_core=10 +nparticles_per_core=100 nparticles = number_of_processors_to_use*nparticles_per_core -R0 = 1.25#jnp.linspace(1.23, 1.27, nparticles) +s=0.25 trace_tolerance = 1e-7 -times_to_trace=100 +times_to_trace=1000 mass=PROTON_MASS mass_a=4.*mass mass_e=ELECTRON_MASS @@ -42,15 +41,15 @@ # Load coils and field -json_file = os.path.join(os.path.dirname(__file__), '..', 'input_files', 'ESSOS_biot_savart_LandremanPaulQA.json') -coils = Coils.from_json(json_file) -field = BiotSavart(coils) +wout_file = os.path.join(os.path.dirname(__file__), '..', 'input_files', "wout_LandremanPaul2021_QA_reactorScale_lowres.nc") +vmec = Vmec(wout_file, ntheta=60, nphi=60, range_torus='half period', close=True) -# Initialize particles -Z0 = jnp.zeros(nparticles) -phi0 = jnp.zeros(nparticles) -initial_xyz=jnp.array([R0*jnp.cos(phi0), R0*jnp.sin(phi0), Z0]).T -particles = Particles(initial_xyz=initial_xyz,initial_vparallel_over_v=-1.*jnp.ones(nparticles), mass=mass, energy=energy) +theta = jnp.zeros(nparticles) +phi = jnp.zeros(nparticles) + +initial_xyz=jnp.array([[s]*nparticles, theta, phi]).T +particles = Particles(initial_xyz=initial_xyz, mass=mass, + charge=ELEMENTARY_CHARGE, energy=energy, field=vmec,initial_vparallel_over_v=1.0*jnp.ones(nparticles)) #Initialize background species @@ -82,67 +81,47 @@ pitch_mean=0. pitch_sigma=jnp.sqrt(2.**2/12) -#import jax -#import jax.numpy as jnp -#from essos.dynamics import GuidingCenterCollisionsDriftMu as GCCD -#from essos.dynamics import GuidingCenterCollisionsDiffusionMu as GCCDiff -#from essos.background_species import nu_s_ab,nu_D_ab,nu_par_ab, d_nu_par_ab -#B_particle=jax.vmap(field.AbsB,in_axes=0)(particles.initial_xyz) -#mu=particles.initial_vperpendicular**2*particles.mass*0.5/B_particle/particles.mass -#initial_conditions = jnp.concatenate([particles.initial_xyz,particles.initial_vparallel[:, None],mu[:, None]],axis=1) -#args = (field, particles,species) -#GCCD(0,initial_conditions[0],args) -#GCCDiff(0,initial_conditions[0],args) -#initial_condition=initial_conditions[0] -#initial_condition = jnp.concatenate([particles.initial_xyz,total_speed_temp[:, None], particles.initial_vparallel_over_v[:, None]], axis=1)[0] -#initial_condition = jnp.concatenate([particles.initial_xyz,total_speed_temp[:, None], particles.initial_vparallel_over_v[:, None]], axis=1)[0] - # Trace in ESSOS time0 = time() -tracing = Tracing(field=field, model='GuidingCenterCollisionsMuFixed', particles=particles, +tracing = Tracing(field=vmec, model='GuidingCenterCollisionsMuFixed', particles=particles, maxtime=tmax, timestep=dt,times_to_trace=times_to_trace,species=species,tag_gc=0.) print(f"ESSOS tracing took {time()-time0:.2f} seconds") trajectories = tracing.trajectories - -# Plot trajectories, velocity parallel to the magnetic field, and energy error fig = plt.figure(figsize=(9, 8)) ax1 = fig.add_subplot(221, projection='3d') ax2 = fig.add_subplot(222) ax3 = fig.add_subplot(223) ax4 = fig.add_subplot(224) -coils.plot(ax=ax1, show=False) -tracing.plot(ax=ax1, show=False) +#vmec.plot(ax=ax1, show=False) +#tracing.plot(ax=ax1, show=False) -v=jnp.sqrt(tracing.energy*2./particles.mass) +# Plot only a random subset of 10 particles in 3D +subset_size = 10 +subset_indices = np.random.choice(len(trajectories), subset_size, replace=False) -for i, trajectory in enumerate(trajectories): - #ax2.plot(tracing.times, (tracing.energy[i]-tracing.energy[i,0])/tracing.energy[i,0], label=f'Particle {i+1}') - ax2.plot(tracing.times, (v[i]-v[i,0])/v[i,0], label=f'Particle {i+1}') - ax3.plot(tracing.times, trajectory[:, 3]/jnp.sqrt(tracing.energy[i]/mass*2.), label=f'Particle {i+1}') +for i in subset_indices: + trajectory = trajectories[i] + ax1.plot(trajectory[:,0], trajectory[:,1], trajectory[:,2], label=f'Particle {i+1}') + ax2.plot(tracing.times, (tracing.energy()[i]-tracing.energy()[i,0])/tracing.energy()[i,0], label=f'Particle {i+1}') + ax3.plot(tracing.times, 299792458*trajectory[:, 3]/jnp.sqrt(tracing.energy()[i]/mass*2.), label=f'Particle {i+1}') ax4.plot(jnp.sqrt(trajectory[:,0]**2+trajectory[:,1]**2), trajectory[:, 2], label=f'Particle {i+1}') - - -ax2.set_xlabel('Time (s)') +ax2.set_xlabel(r'$t~[\mathrm{s}]$') ax2.set_ylabel('Normalized energy variation') ax3.set_ylabel(r'$v_{\parallel}/v$') -#ax2.legend() -ax3.set_xlabel('Time (s)') -#ax3.legend() -ax4.set_xlabel('R (m)') -ax4.set_ylabel('Z (m)') -#ax4.legend() +ax3.set_xlabel(r'$t~[\mathrm{s}]$') +ax4.set_xlabel(r'$R~[\mathrm{m}]$') +ax4.set_ylabel(r'$Z~[\mathrm{m}]$') plt.tight_layout() -plt.savefig('traj.pdf') +plt.savefig('traj_time.pdf') - -v=jnp.sqrt(tracing.energy*2./particles.mass) -#pitch=trajectories[:,:,3]/v -vpar=trajectories[:,:,3] -vperp=tracing.vperp_final +v=jnp.sqrt(tracing.energy()*2./particles.mass) +vpar=trajectories[:,:,3]*SPEED_OF_LIGHT +vpar=jnp.where(jnp.isfinite(vpar), vpar, jnp.nan) +vperp=tracing.v_perp() pitch=vpar/v # Plot distribution in velocities initial t and final fig3 = plt.figure(figsize=(9, 8)) @@ -198,7 +177,6 @@ plt.savefig('statistics.pdf') - # Plot distribution in velocities initial t and final fig2 = plt.figure(figsize=(9, 8)) ax12 = fig2.add_subplot(251) @@ -212,12 +190,12 @@ ax92 = fig2.add_subplot(259) nbins=64 -v0=jnp.sqrt(tracing.energy[:,0]*2./particles.mass) -vfinal=jnp.sqrt(tracing.energy[:,-1]*2./particles.mass) -vperp0=tracing.vperp_final[:,0] -vperpfinal=tracing.vperp_final[:,-1] -vpar0=trajectories[:,0,3] -vparfinal=trajectories[:,-1,3] +v0=jnp.sqrt(tracing.energy()[:,0]*2./particles.mass)/SPEED_OF_LIGHT +vfinal=jnp.sqrt(tracing.energy()[:,-1]*2./particles.mass)/SPEED_OF_LIGHT +vperp0=tracing.v_perp()[:,0]/SPEED_OF_LIGHT +vperpfinal=tracing.v_perp()[:,-1]/SPEED_OF_LIGHT +vpar0=vpar[:,0]/SPEED_OF_LIGHT +vparfinal=vpar[:,-1]/SPEED_OF_LIGHT pitch0=vpar0/v0 pitch_final=vparfinal/vfinal @@ -339,7 +317,16 @@ def find_first_less_than_numpy(arr, value): ax92.set_xlabel(r'$t_{final}$') plt.tight_layout() -plt.savefig('dist.pdf') + +# Improved time distribution plot +plt.figure(figsize=(7, 5)) +plt.hist(good_t_final, bins=nbins, color='c', edgecolor='black', alpha=0.7) +plt.title(r'$t_{final}$ Distribution', fontweight='bold') +plt.xlabel(r'$t_{final}$', fontweight='bold') +plt.ylabel('Counts', fontweight='bold') +plt.tight_layout() +plt.rcParams.update({'font.size': 18, 'font.weight': 'bold'}) +plt.savefig('dist_time.pdf', dpi=300) ## Save results in vtk format to analyze in Paraview # tracing.to_vtk('trajectories') diff --git a/examples/particle_tracing_collisions/trace_particles_coils_guidingcenter_with_classifier_with_collisionsMu.py b/examples/particle_tracing_collisions/trace_particles_coils_guidingcenter_with_classifier_with_collisionsMu.py index 4535be1..95a01cf 100644 --- a/examples/particle_tracing_collisions/trace_particles_coils_guidingcenter_with_classifier_with_collisionsMu.py +++ b/examples/particle_tracing_collisions/trace_particles_coils_guidingcenter_with_classifier_with_collisionsMu.py @@ -12,7 +12,6 @@ from essos.dynamics import Tracing, Particles from essos.background_species import BackgroundSpecies - # Input parameters tmax = 1e-4 timestep=1.e-8 @@ -24,7 +23,6 @@ rtol=1.e-5 energy=FUSION_ALPHA_PARTICLE_ENERGY - #Initialize background species #number_species=2 #(electrons,deuterium) #mass_array=jnp.array([ELECTRON_MASS/PROTON_MASS,2]) #mass_over_mproton @@ -89,8 +87,8 @@ for i, trajectory in enumerate(trajectories): #ax2.plot(tracing.times, jnp.abs(tracing.energy[i]-particles.energy)/particles.energy, label=f'Particle {i+1}') - ax2.plot(tracing.times, (tracing.energy[i]-tracing.energy[i][0])/particles.energy, label=f'Particle {i+1}') - ax3.plot(tracing.times, trajectory[:, 3]*SPEED_OF_LIGHT/particles.total_speed, label=f'Particle {i+1}') + ax2.plot(tracing.times, (tracing.energy()[i]-tracing.energy()[i][0])/tracing.energy()[i][0], label=f'Particle {i+1}') + ax3.plot(tracing.times, trajectory[:, 3]*SPEED_OF_LIGHT/tracing.energy()[i], label=f'Particle {i+1}') #ax4.plot(jnp.sqrt(trajectory[:,0]**2+trajectory[:,1]**2), trajectory[:, 2], label=f'Particle {i+1}') ax4.plot(jnp.sqrt(trajectory[:,0]**2+trajectory[:,1]**2), trajectory[:, 2], label=f'Particle {i+1}') ax2.set_xlabel('Time (s)') diff --git a/examples/particle_tracing_collisions/trace_particles_vmec_collisionsMu.py b/examples/particle_tracing_collisions/trace_particles_vmec_collisionsMu.py index 235aeaa..a5edb6d 100644 --- a/examples/particle_tracing_collisions/trace_particles_vmec_collisionsMu.py +++ b/examples/particle_tracing_collisions/trace_particles_vmec_collisionsMu.py @@ -10,7 +10,6 @@ from essos.background_species import BackgroundSpecies import numpy as np - # Input parameters tmax = 1.e-4 timestep=1.e-8 @@ -71,9 +70,9 @@ for i in np.random.choice(nparticles, size=n_particles_to_plot, replace=False): trajectory = trajectories[i] ## Plot energy error - ax2.plot(tracing.times, (tracing.energy[i]-tracing.energy[i][0])/tracing.energy[i][0], label=f'Particle {i+1}') + ax2.plot(tracing.times, (tracing.energy()[i]-tracing.energy()[i,0])/tracing.energy()[i,0], label=f'Particle {i+1}') ## Plot velocity parallel to the magnetic field - ax3.plot(tracing.times, trajectory[:, 3]*SPEED_OF_LIGHT/jnp.sqrt(tracing.energy[i]/particles.mass*2.), label=f'Particle {i+1}') + ax3.plot(tracing.times, trajectory[:, 3]*SPEED_OF_LIGHT/jnp.sqrt(tracing.energy()[i]/particles.mass*2.), label=f'Particle {i+1}') ## Plot s-coordinate ax4.plot(tracing.times, trajectory[:,0], label=f'Particle {i+1}') # ax4.set_ylabel(r'$s=\psi/\psi_b$') diff --git a/examples/coils_from_BOOZ_XFORM.py b/examples/simple_examples/coils_from_BOOZ_XFORM.py similarity index 100% rename from examples/coils_from_BOOZ_XFORM.py rename to examples/simple_examples/coils_from_BOOZ_XFORM.py diff --git a/examples/coils_from_nearaxis.py b/examples/simple_examples/coils_from_nearaxis.py similarity index 100% rename from examples/coils_from_nearaxis.py rename to examples/simple_examples/coils_from_nearaxis.py diff --git a/examples/simple_examples/create_perturbed_coils.py b/examples/simple_examples/create_perturbed_coils.py index 6368366..f0e8ee1 100644 --- a/examples/simple_examples/create_perturbed_coils.py +++ b/examples/simple_examples/create_perturbed_coils.py @@ -12,9 +12,6 @@ from functools import partial from essos.coil_perturbation import GaussianSampler, perturb_curves - - - # Coils parameters order_Fourier_series_coils = 4 number_coil_points = 80 @@ -60,8 +57,6 @@ plt.legend() plt.show() - - # # Save the coils to a json file # coils_optimized.to_json("stellarator_coils.json") # # Load the coils from a json file diff --git a/examples/get_derivatives_coils_particle_confinement_guidingcenter.py b/examples/simple_examples/get_derivatives_coils_particle_confinement_guidingcenter.py similarity index 100% rename from examples/get_derivatives_coils_particle_confinement_guidingcenter.py rename to examples/simple_examples/get_derivatives_coils_particle_confinement_guidingcenter.py diff --git a/tests/test_coil_perturbation.py b/tests/test_coil_perturbation.py index a33821e..4405ddb 100644 --- a/tests/test_coil_perturbation.py +++ b/tests/test_coil_perturbation.py @@ -108,7 +108,7 @@ def test_perturb_curves_systematic(self): key = jax.random.PRNGKey(0) for sampler in [sampler0, sampler1, sampler2]: curves = DummyCurves(n_base_curves=2, nfp=1, stellsym=True, n_points=5) - perturb_curves_systematic(curves, sampler, key) + curves = perturb_curves_systematic(curves, sampler, key) # Just check that gamma arrays are still the right shape self.assertEqual(curves.gamma.shape, (2, 5, 3)) @@ -120,8 +120,8 @@ def test_perturb_curves_statistic(self): key = jax.random.PRNGKey(0) for sampler in [sampler0, sampler1, sampler2]: curves = DummyCurves(n_base_curves=2, nfp=1, stellsym=True, n_points=5) - perturb_curves_statistic(curves, sampler, key) + curves = perturb_curves_statistic(curves, sampler, key) self.assertEqual(curves.gamma.shape, (2, 5, 3)) if __name__ == "__main__": - unittest.main() \ No newline at end of file + unittest.main() diff --git a/tests/test_objective_functions.py b/tests/test_objective_functions.py index 857e8fe..0bec838 100644 --- a/tests/test_objective_functions.py +++ b/tests/test_objective_functions.py @@ -1,260 +1,275 @@ import unittest +from types import SimpleNamespace from unittest.mock import MagicMock, patch + +import jax import jax.numpy as jnp import essos.objective_functions as objf + class DummyField: def __init__(self): - self.R0 = jnp.array([1.]) - self.Z0 = jnp.array([0.]) - self.phi = jnp.array([0.]) - self.B_axis = jnp.array([[1., 0., 0.]]) - self.grad_B_axis = jnp.array([[0., 0., 0.]]) + 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.AbsB = MagicMock(return_value=5.7) - self.B = MagicMock(return_value=jnp.array([1., 0., 0.])) - self.dB_by_dX = MagicMock(return_value=jnp.array([0., 0., 0.])) - self.B_covariant = MagicMock(return_value=jnp.array([1., 0., 0.])) - self.coils_length = jnp.array([30.]) - self.coils_curvature = jnp.ones((2, 10)) - self.gamma = jnp.zeros((2, 10, 3)) - self.gamma_dash = jnp.ones((2, 10, 3)) - self.gamma_dashdash = jnp.ones((2, 10, 3)) - self.currents = jnp.ones(2) - self.quadpoints = jnp.linspace(0, 1, 10) - self.x = jnp.zeros((10,1)) - -class DummyCoils(DummyField): - def __init__(self): - super().__init__() + self.coils = DummyCoils() + + def AbsB(self, points): + return 5.7 + 0.1 * jnp.sum(points) + + def B(self, points): + return jnp.array([1.0 + 0.1 * points[0], 0.5, 0.25]) + + def dB_by_dX(self, points): + return jnp.eye(3) + + def B_covariant(self, points): + return jnp.array([1.0, 0.5, 0.25]) + + def copy(self): + return DummyField() -class DummyCurves: - def __init__(self, *args, **kwargs): - pass class DummyParticles: def __init__(self): - self.to_full_orbit = MagicMock() - self.trajectories = jnp.zeros((2, 10, 3)) + self.energy = 1.0 + self.mass = 1.0 + self.charge = 1.0 + self.total_speed = 1.0 + + def to_full_orbit(self, field): + return None + class DummyTracing: def __init__(self, *args, **kwargs): - self.trajectories = jnp.zeros((2, 10, 3)) - self.field = DummyField() - self.loss_fractions = jnp.array([0.1,0.2,1.]) - self.times_to_trace = 10 + xyz = jnp.array( + [ + [[1.0, 0.0, 0.0], [1.1, 0.0, 0.05], [1.2, 0.0, 0.1], [1.3, 0.0, 0.15]], + [[1.0, 0.0, 0.0], [1.05, 0.0, 0.04], [1.1, 0.0, 0.08], [1.15, 0.0, 0.12]], + ], + dtype=jnp.float64, + ) + self.trajectories = xyz + self.field = kwargs.get("field", DummyField()) + self.loss_fractions = jnp.array([0.1, 0.2, 1.0]) + self.times_to_trace = 4 self.maxtime = 1e-5 -class DummyVmec: - def __init__(self): - self.surface = MagicMock() class DummySurface: def __init__(self): - self.gamma = jnp.zeros((10, 3)) - self.unitnormal = jnp.ones((10, 3)) - -def dummy_sampler(*args, **kwargs): - return 0 - -def dummy_new_nearaxis_from_x_and_old_nearaxis(x, field_nearaxis): - class DummyNearAxis: - elongation = jnp.array([1.]) - iota = 1.0 - x = jnp.array([1.]) - R0 = jnp.array([1.]) - Z0 = jnp.array([0.]) - phi = jnp.array([0.]) - B_axis = jnp.array([[1., 0., 0.]]) - grad_B_axis = jnp.array([[0., 0., 0.]]) - return DummyNearAxis() + 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 + + +@jax.tree_util.register_pytree_node_class +class PytreeCoils: + def __init__(self, gamma, gamma_dash, gamma_dashdash, currents, quadpoints, length, curvature, base_curves, order=1, nfp=1, stellsym=False): + self.gamma = gamma + self.gamma_dash = gamma_dash + self.gamma_dashdash = gamma_dashdash + self.currents = currents + self.length = length + self.curvature = curvature + self.order = order + self.nfp = nfp + self.stellsym = stellsym + self.curves = SimpleNamespace(quadpoints=quadpoints, curves=base_curves) + + def __len__(self): + return self.gamma.shape[0] + + def copy(self): + return PytreeCoils( + self.gamma, + self.gamma_dash, + self.gamma_dashdash, + self.currents, + self.curves.quadpoints, + self.length, + self.curvature, + self.curves.curves, + order=self.order, + nfp=self.nfp, + stellsym=self.stellsym, + ) + + def tree_flatten(self): + children = ( + self.gamma, + self.gamma_dash, + self.gamma_dashdash, + self.currents, + self.curves.quadpoints, + self.length, + self.curvature, + self.curves.curves, + ) + aux = {"order": self.order, "nfp": self.nfp, "stellsym": self.stellsym} + return children, aux + + @classmethod + def tree_unflatten(cls, aux, children): + return cls(*children, **aux) + + +@jax.tree_util.register_pytree_node_class +class PytreeSurface: + def __init__(self, gamma, unitnormal, stellsym=False, nfp=1): + self.gamma = gamma + self.unitnormal = 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, unitnormal, stellsym=stellsym, nfp=nfp) + class TestObjectiveFunctions(unittest.TestCase): def setUp(self): - self.x = jnp.ones(12) - self.dofs_curves = jnp.ones((2, 3)) - self.currents_scale = 1.0 - self.nfp = 1 - self.n_segments = 10 - self.stellsym = True - self.key = 0 - self.sampler = dummy_sampler self.field = DummyField() - self.coils = DummyCoils() - self.curves = DummyCurves() + self.field_nearaxis = DummyField() self.particles = DummyParticles() - self.tracing = DummyTracing() - self.vmec = DummyVmec() self.surface = DummySurface() - - @patch('essos.objective_functions.Curves', return_value=DummyCurves()) - @patch('essos.objective_functions.Coils', return_value=DummyCoils()) - @patch('essos.objective_functions.BiotSavart', return_value=DummyField()) - @patch('essos.objective_functions.perturb_curves', side_effect=lambda curves, sampler, key=None, perturbation_type=None: curves) - def test_perturbed_field_and_coils_from_dofs(self, pc, bs, coils, curves): - objf.pertubred_field_from_dofs(self.x, self.key, self.sampler, self.dofs_curves, self.currents_scale, self.nfp) - objf.perturbed_coils_from_dofs(self.x, self.key, self.sampler, self.dofs_curves, self.currents_scale, self.nfp) - - @patch('essos.objective_functions.Curves', return_value=DummyCurves()) - @patch('essos.objective_functions.Coils', return_value=DummyCoils()) - @patch('essos.objective_functions.BiotSavart', return_value=DummyField()) - def test_field_and_coils_from_dofs(self, bs, coils, curves): - objf.field_from_dofs(self.x, self.dofs_curves, self.currents_scale, self.nfp) - objf.coils_from_dofs(self.x, self.dofs_curves, self.currents_scale, self.nfp) - objf.curves_from_dofs(self.x, self.dofs_curves, self.nfp) - - @patch('essos.objective_functions.field_from_dofs', return_value=DummyField()) - def test_loss_coil_length_and_curvature(self, ffd): - objf.loss_coil_length(self.x, self.dofs_curves, self.currents_scale, self.nfp) - objf.loss_coil_curvature(self.x, self.dofs_curves, self.currents_scale, self.nfp) - objf.loss_coil_length_new(self.x, self.dofs_curves, self.currents_scale, self.nfp) - objf.loss_coil_curvature_new(self.x, self.dofs_curves, self.currents_scale, self.nfp) - - @patch('essos.objective_functions.field_from_dofs', return_value=DummyField()) - def test_loss_normB_axis(self, ffd): - objf.loss_normB_axis(self.x, self.dofs_curves, self.currents_scale, self.nfp) - objf.loss_normB_axis_average(self.x, self.dofs_curves, self.currents_scale, self.nfp) - - @patch('essos.objective_functions.field_from_dofs', return_value=DummyField()) - def test_loss_particle_functions(self, ffd): - with patch('essos.objective_functions.Tracing', return_value=self.tracing): - objf.loss_particle_radial_drift(self.x, self.particles, self.dofs_curves, self.currents_scale, self.nfp) - objf.loss_particle_alpha_drift(self.x, self.particles, self.dofs_curves, self.currents_scale, self.nfp) - objf.loss_particle_gamma_c(self.x, self.particles, self.dofs_curves, self.currents_scale, self.nfp) - objf.loss_particle_r_cross_final(self.x, self.particles, self.dofs_curves, self.currents_scale, self.nfp) - objf.loss_particle_r_cross_max_constraint(self.x, self.particles, self.dofs_curves, self.currents_scale, self.nfp) - objf.loss_Br(self.x, self.particles, self.dofs_curves, self.currents_scale, self.nfp) - objf.loss_iota(self.x, self.particles, self.dofs_curves, self.currents_scale, self.nfp) - - @patch('essos.objective_functions.field_from_dofs', return_value=DummyField()) - def test_loss_lost_fraction(self, ffd): - with patch('essos.objective_functions.Tracing', return_value=self.tracing): - objf.loss_lost_fraction(self.field, self.particles, self.dofs_curves, self.currents_scale, self.nfp) - - def test_normB_axis(self): - objf.normB_axis(self.field) - - @patch('essos.objective_functions.field_from_dofs', return_value=DummyField()) - @patch('essos.objective_functions.new_nearaxis_from_x_and_old_nearaxis', side_effect=dummy_new_nearaxis_from_x_and_old_nearaxis) - def test_loss_coils_for_nearaxis_and_loss_coils_and_nearaxis(self, nna, ffd): - objf.loss_coils_for_nearaxis(self.x, self.field, self.dofs_curves, self.currents_scale, self.nfp) - objf.loss_coils_and_nearaxis(jnp.ones(13), self.field, self.dofs_curves, self.currents_scale, self.nfp) - - def test_difference_B_gradB_onaxis(self): - objf.difference_B_gradB_onaxis(self.field, self.field) - - @patch('essos.objective_functions.Curves', return_value=DummyCurves()) - @patch('essos.objective_functions.Coils', return_value=DummyCoils()) - @patch('essos.objective_functions.BiotSavart', return_value=DummyField()) - @patch('essos.objective_functions.BdotN_over_B', return_value=jnp.ones(10)) - def test_loss_bdotn_over_b(self, bdotn, bs, coils, curves): - objf.loss_bdotn_over_b(self.x, self.vmec, self.dofs_curves, self.currents_scale, self.nfp) - - @patch('essos.objective_functions.field_from_dofs', return_value=DummyField()) - @patch('essos.objective_functions.BdotN_over_B', return_value=jnp.ones(10)) - def test_loss_BdotN(self, bdotn, ffd): - objf.loss_BdotN(self.x, self.vmec, self.dofs_curves, self.currents_scale, self.nfp) - - @patch('essos.objective_functions.field_from_dofs', return_value=DummyField()) - @patch('essos.objective_functions.BdotN_over_B', return_value=jnp.ones(10)) - def test_loss_BdotN_only(self, bdotn, ffd): - objf.loss_BdotN_only(self.x, self.vmec, self.dofs_curves, self.currents_scale, self.nfp) - - @patch('essos.objective_functions.field_from_dofs', return_value=DummyField()) - @patch('essos.objective_functions.BdotN_over_B', return_value=jnp.ones(10)) - def test_loss_BdotN_only_constraint(self, bdotn, ffd): - objf.loss_BdotN_only_constraint(self.x, self.vmec, self.dofs_curves, self.currents_scale, self.nfp) - - @patch('essos.objective_functions.BdotN_over_B', return_value=jnp.ones(10)) - @patch('essos.objective_functions.pertubred_field_from_dofs', return_value=DummyField()) - def test_loss_BdotN_only_stochastic(self, perturbed, bdotn): - objf.loss_BdotN_only_stochastic(self.x, self.sampler, 2, self.vmec, self.dofs_curves, self.currents_scale, self.nfp) - - @patch('essos.objective_functions.BdotN_over_B', return_value=jnp.ones(10)) - @patch('essos.objective_functions.pertubred_field_from_dofs', return_value=DummyField()) - def test_loss_BdotN_only_constraint_stochastic(self, perturbed, bdotn): - objf.loss_BdotN_only_constraint_stochastic(self.x, self.sampler, 2, self.vmec, self.dofs_curves, self.currents_scale, self.nfp) - - @patch('essos.objective_functions.coils_from_dofs', return_value=DummyCoils()) - def test_loss_cs_distance_and_array(self, cfd): - objf.loss_cs_distance(self.x, self.surface, self.dofs_curves, self.currents_scale, self.nfp) - objf.loss_cs_distance_array(self.x, self.surface, self.dofs_curves, self.currents_scale, self.nfp) - - @patch('essos.objective_functions.coils_from_dofs', return_value=DummyCoils()) - def test_loss_cc_distance_and_array(self, cfd): - objf.loss_cc_distance(self.x, self.dofs_curves, self.currents_scale, self.nfp) - objf.loss_cc_distance_array(self.x, self.dofs_curves, self.currents_scale, self.nfp) - - @patch('essos.objective_functions.coils_from_dofs', return_value=DummyCoils()) - def test_loss_linking_mnumber_and_constraint(self, cfd): - objf.loss_linking_mnumber(self.x, self.dofs_curves, self.currents_scale, self.nfp) - objf.loss_linking_mnumber_constarint(self.x, self.dofs_curves, self.currents_scale, self.nfp) - - def test_cc_distance_pure(self): - gamma1 = jnp.ones((10, 3))*3. - l1 = jnp.ones((10, 3)) - gamma2 = jnp.ones((10, 3))*4. - l2 = jnp.ones((10, 3))*6. - objf.cc_distance_pure(gamma1, l1, gamma2, l2, 1.0) - - def test_cs_distance_pure(self): - gammac = jnp.ones((10, 3))*7. - lc = jnp.ones((10, 3)) - gammas = jnp.ones((10, 3))*9. - ns = jnp.ones((10, 3))*10. - objf.cs_distance_pure(gammac, lc, gammas, ns, 1.0) - - @patch('essos.objective_functions.coils_from_dofs', return_value=DummyCoils()) - def test_loss_lorentz_force_coils(self, cfd): - objf.loss_lorentz_force_coils(self.x, self.dofs_curves, self.currents_scale, self.nfp) - - @patch('essos.objective_functions.compute_curvature', return_value=1.0) - @patch('essos.objective_functions.BiotSavart_from_gamma', return_value=MagicMock(B=MagicMock(return_value=jnp.array([1., 0., 0.])))) - def test_lp_force_pure(self, bsg, cc): - gamma = jnp.ones((2, 10, 3))*2. - gamma_dash = jnp.ones((2, 10, 3))*3. - gamma_dashdash = jnp.ones((2, 10, 3)) - currents = jnp.ones(2) - quadpoints = jnp.linspace(0, 1, 10) - objf.lp_force_pure(0, gamma, gamma_dash, gamma_dashdash, currents, quadpoints, 1, 1e6) - - def test_B_regularized_singularity_term(self): + self.sampler = MagicMock(name="sampler") + self.keys = jnp.array([0, 1], dtype=jnp.int32) + self.coils = PytreeCoils( + gamma=jnp.arange(2 * 5 * 3, dtype=jnp.float64).reshape(2, 5, 3) / 10.0, + gamma_dash=jnp.ones((2, 5, 3), dtype=jnp.float64), + gamma_dashdash=jnp.ones((2, 5, 3), dtype=jnp.float64) * 0.1, + currents=jnp.ones(2, dtype=jnp.float64), + quadpoints=jnp.linspace(0.0, 1.0, 5), + length=jnp.array([3.0, 4.0], dtype=jnp.float64), + curvature=jnp.ones((2, 5), dtype=jnp.float64), + base_curves=jnp.arange(2 * 3 * 3, dtype=jnp.float64).reshape(2, 3, 3) / 10.0, + ) + self.pytree_surface = PytreeSurface( + gamma=jnp.arange(2 * 3 * 3, dtype=jnp.float64).reshape(2, 3, 3) / 20.0, + unitnormal=jnp.ones((2, 3, 3), dtype=jnp.float64), + ) + + 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.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))) + self.assertTrue(jnp.isfinite(objf.loss_r0_near_axis(self.field_nearaxis))) + + @patch("essos.objective_functions.Tracing", side_effect=DummyTracing) + def test_particle_losses(self, tracing): + self.assertTrue(jnp.isfinite(objf.loss_particle_radial_drift(self.field, self.particles))) + self.assertTrue(jnp.isfinite(objf.loss_particle_radial_drift_fullorbit(self.field, self.particles))) + self.assertTrue(jnp.isfinite(objf.loss_particle_alpha_drift(self.field, self.particles))) + self.assertTrue(jnp.isfinite(objf.loss_particle_gammac(self.field, self.particles))) + 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))) + + @patch("essos.objective_functions.BdotN_over_B", return_value=jnp.ones((2, 3), dtype=jnp.float64)) + def test_surface_losses(self, bdotn): + self.assertTrue(jnp.isfinite(objf.normB_axis(self.field)).all()) + self.assertTrue(jnp.isfinite(objf.loss_normB_axis_average(self.field))) + self.assertTrue(jnp.isfinite(objf.loss_BdotN(self.field, self.surface))) + self.assertTrue(jnp.isfinite(objf.loss_BdotN_constraint(self.field, self.surface))) + + def test_copy_coils_from_field(self): + copied = objf.copy_coils_from_field(self.field) + self.assertIsInstance(copied, DummyCoils) + + @patch("essos.objective_functions.BiotSavart", return_value=DummyField()) + @patch("essos.objective_functions.perturb_curves_systematic") + @patch("essos.objective_functions.perturb_curves_statistic") + def test_perturbed_field_from_field(self, statistical, systematic, biot_savart): + systematic.side_effect = lambda coils, sampler, key=None: coils + statistical.side_effect = lambda coils, sampler, key=None: coils + field = objf.perturbed_field_from_field(self.field, 0, self.sampler) + self.assertIsInstance(field, DummyField) + systematic.assert_called_once() + statistical.assert_called_once() + + @patch("essos.objective_functions.BdotN_over_B", return_value=jnp.ones((2, 3), dtype=jnp.float64)) + @patch("essos.objective_functions.perturbed_field_from_field", return_value=DummyField()) + def test_stochastic_surface_losses(self, perturbed_field, bdotn): + self.assertTrue(jnp.isfinite(objf.loss_bdotn_stochastic(self.field, self.surface, self.sampler, self.keys))) + self.assertTrue(jnp.isfinite(objf.constraint_bdotn_stochastic(self.field, self.surface, self.sampler, self.keys))) + + def test_coil_length_and_curvature_losses(self): + self.assertTrue(jnp.isfinite(objf.loss_coil_length(self.coils, max_coil_length=4.0)).all()) + self.assertTrue(jnp.isfinite(objf.loss_coil_curvature(self.coils, max_coil_curvature=2.0)).all()) + + def test_compute_candidates(self): + i_vals, j_vals = objf.compute_candidates(self.coils, min_separation=10.0) + self.assertEqual(i_vals.ndim, 1) + self.assertEqual(j_vals.ndim, 1) + + def test_blockwise_losses_with_non_divisible_blocks(self): + separation = objf.loss_coil_separation(self.coils, 0.5, block_size=3) + surface_distance = objf.loss_coil_surface_distance(self.coils, self.pytree_surface, 0.5, block_size=4) + linking = objf.loss_linkingnumber(self.coils, block_size=4) + self.assertTrue(jnp.isfinite(separation)) + self.assertTrue(jnp.isfinite(surface_distance)) + self.assertTrue(jnp.isfinite(linking)) + self.assertAlmostEqual(float(separation), float(objf.loss_coil_separation.__wrapped__(self.coils, 0.5, block_size=3))) + self.assertAlmostEqual(float(surface_distance), float(objf.loss_coil_surface_distance.__wrapped__(self.coils, self.pytree_surface, 0.5, block_size=4))) + self.assertAlmostEqual(float(linking), float(objf.loss_linkingnumber.__wrapped__(self.coils, block_size=4))) + + @patch("essos.objective_functions.Curves.compute_curvature", return_value=jnp.ones(5)) + @patch("essos.objective_functions.BiotSavart_from_gamma") + def test_loss_lorentz_force_coils(self, biot_savart_from_gamma, compute_curvature): + class DummyBS: + def B(self, point): + return jnp.zeros(3) + + biot_savart_from_gamma.return_value = DummyBS() + force_loss = objf.loss_lorentz_force_coils(self.coils, threshold=1e6, block_size=2) + self.assertTrue(jnp.isfinite(force_loss)) + self.assertAlmostEqual( + float(force_loss), + float(objf.loss_lorentz_force_coils.__wrapped__(self.coils, threshold=1e6, block_size=2)), + ) + + def test_regularization_helpers(self): rc_prime = jnp.ones((10, 3)) rc_prime_prime = jnp.ones((10, 3)) - objf.B_regularized_singularity_term(rc_prime, rc_prime_prime, 1.0) - - def test_B_regularized_pure(self): - gamma = jnp.ones((10, 3))*4. + gamma = jnp.ones((10, 3)) * 4.0 gammadash = jnp.ones((10, 3)) gammadashdash = jnp.ones((10, 3)) quadpoints = jnp.linspace(0, 1, 10) - current = 1.0 - regularization = 1.0 - objf.B_regularized_pure(gamma, gammadash, gammadashdash, quadpoints, current, regularization) - - def test_regularization_circ(self): + self.assertTrue(jnp.isfinite(objf.B_regularized_singularity_term(rc_prime, rc_prime_prime, 1.0)).all()) + self.assertTrue(jnp.isfinite(objf.B_regularized_pure(gamma, gammadash, gammadashdash, quadpoints, 1.0, 1.0)).all()) self.assertTrue(objf.regularization_circ(2.0) > 0) + self.assertTrue(jnp.isfinite(objf.regularization_rect(2.0, 1.0))) + self.assertTrue(jnp.isfinite(objf.rectangular_xsection_k(2.0, 1.0))) + 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() - def test_regularization_rect_and_k_and_delta(self): - a, b = 2.0, 1.0 - objf.regularization_rect(a, b) - objf.rectangular_xsection_k(a, b) - objf.rectangular_xsection_delta(a, b) - - def test_linking_number_pure_and_integrand(self): - gamma1 = jnp.ones((10, 3))*4. - lc1 = jnp.ones((10, 3))*2. - gamma2 = jnp.ones((10, 3))*6. - lc2 = jnp.ones((10, 3))*5. - dphi = 0.1 - objf.linking_number_pure(gamma1, lc1, gamma2, lc2, dphi) - r1 = jnp.zeros(3) - dr1 = jnp.ones(3) - r2 = jnp.zeros(3) - dr2 = jnp.ones(3) - objf.integrand_linking_number(r1, dr1, r2, dr2, dphi, dphi) if __name__ == "__main__": unittest.main()