essos 0.1__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- essos/__init__.py +0 -0
- essos/__main__.py +20 -0
- essos/coils.py +437 -0
- essos/constants.py +7 -0
- essos/dynamics.py +411 -0
- essos/fields.py +470 -0
- essos/objective_functions.py +170 -0
- essos/optimization.py +81 -0
- essos/plot.py +21 -0
- essos/surfaces.py +183 -0
- essos/version.py +21 -0
- essos-0.1.dist-info/LICENSE +21 -0
- essos-0.1.dist-info/METADATA +248 -0
- essos-0.1.dist-info/RECORD +17 -0
- essos-0.1.dist-info/WHEEL +5 -0
- essos-0.1.dist-info/entry_points.txt +2 -0
- essos-0.1.dist-info/top_level.txt +1 -0
essos/__init__.py
ADDED
|
File without changes
|
essos/__main__.py
ADDED
|
@@ -0,0 +1,20 @@
|
|
|
1
|
+
"""Main command line interface to ESSOS."""
|
|
2
|
+
import sys
|
|
3
|
+
import tomllib
|
|
4
|
+
|
|
5
|
+
def main(cl_args=sys.argv[1:]):
|
|
6
|
+
"""Run the main ESSOS code from the command line.
|
|
7
|
+
|
|
8
|
+
Reads and parses user input from command line, runs the code,
|
|
9
|
+
and prints and plots the resulting simulation.
|
|
10
|
+
|
|
11
|
+
"""
|
|
12
|
+
if len(cl_args) == 0:
|
|
13
|
+
print("Using standard input parameters instead of an input TOML file.")
|
|
14
|
+
output = 0
|
|
15
|
+
else:
|
|
16
|
+
parameters = tomllib.load(open(cl_args[0], "rb"))
|
|
17
|
+
output = 0
|
|
18
|
+
|
|
19
|
+
if __name__ == "__main__":
|
|
20
|
+
main(sys.argv[1:])
|
essos/coils.py
ADDED
|
@@ -0,0 +1,437 @@
|
|
|
1
|
+
import jax
|
|
2
|
+
jax.config.update("jax_enable_x64", True)
|
|
3
|
+
import jax.numpy as jnp
|
|
4
|
+
from jax.lax import fori_loop
|
|
5
|
+
from jax import tree_util, jit, vmap
|
|
6
|
+
from functools import partial
|
|
7
|
+
from .plot import fix_matplotlib_3d
|
|
8
|
+
|
|
9
|
+
def compute_curvature(gammadash, gammadashdash):
|
|
10
|
+
return jnp.linalg.norm(jnp.cross(gammadash, gammadashdash, axis=1), axis=1) / jnp.linalg.norm(gammadash, axis=1)**3
|
|
11
|
+
|
|
12
|
+
class Curves:
|
|
13
|
+
"""
|
|
14
|
+
Class to store the curves
|
|
15
|
+
|
|
16
|
+
-----------
|
|
17
|
+
Attributes:
|
|
18
|
+
dofs (jnp.ndarray - shape (n_indcurves, 3, 2*order+1)): Fourier Coefficients of the independent curves
|
|
19
|
+
n_segments (int): Number of segments to discretize the curves
|
|
20
|
+
nfp (int): Number of field periods
|
|
21
|
+
stellsym (bool): Stellarator symmetry
|
|
22
|
+
order (int): Order of the Fourier series
|
|
23
|
+
curves jnp.ndarray - shape (n_indcurves*nfp*(1+stellsym), 3, 2*order+1)): Curves obtained by applying rotations and flipping corresponding to nfp fold rotational symmetry and optionally stellarator symmetry
|
|
24
|
+
gamma (jnp.array - shape (n_coils, n_segments, 3)): Discretized curves
|
|
25
|
+
gamma_dash (jnp.array - shape (n_coils, n_segments, 3)): Discretized curves derivatives
|
|
26
|
+
|
|
27
|
+
"""
|
|
28
|
+
def __init__(self, dofs: jnp.ndarray, n_segments: int = 100, nfp: int = 1, stellsym: bool = True):
|
|
29
|
+
dofs = jnp.array(dofs)
|
|
30
|
+
# assert isinstance(dofs, jnp.ndarray), "dofs must be a jnp.ndarray"
|
|
31
|
+
assert dofs.ndim == 3, "dofs must be a 3D array with shape (n_curves, 3, 2*order+1)"
|
|
32
|
+
assert dofs.shape[1] == 3, "dofs must have shape (n_curves, 3, 2*order+1)"
|
|
33
|
+
assert dofs.shape[2] % 2 == 1, "dofs must have shape (n_curves, 3, 2*order+1)"
|
|
34
|
+
assert isinstance(n_segments, int), "n_segments must be an integer"
|
|
35
|
+
assert n_segments > 2, "n_segments must be greater than 2"
|
|
36
|
+
assert isinstance(nfp, int), "nfp must be a positive integer"
|
|
37
|
+
assert nfp > 0, "nfp must be a positive integer"
|
|
38
|
+
assert isinstance(stellsym, bool), "stellsym must be a boolean"
|
|
39
|
+
|
|
40
|
+
self._dofs = dofs
|
|
41
|
+
self._n_segments = n_segments
|
|
42
|
+
self._nfp = nfp
|
|
43
|
+
self._stellsym = stellsym
|
|
44
|
+
self._order = dofs.shape[2]//2
|
|
45
|
+
self._curves = apply_symmetries_to_curves(self.dofs, self.nfp, self.stellsym)
|
|
46
|
+
self.quadpoints = jnp.linspace(0, 1, self.n_segments, endpoint=False)
|
|
47
|
+
self._set_gamma()
|
|
48
|
+
|
|
49
|
+
def __str__(self):
|
|
50
|
+
return f"nfp stellsym order\n{self.nfp} {self.stellsym} {self.order}\n"\
|
|
51
|
+
+ f"Degrees of freedom\n{repr(self.dofs.tolist())}\n"
|
|
52
|
+
|
|
53
|
+
def __repr__(self):
|
|
54
|
+
return f"nfp stellsym order\n{self.nfp} {self.stellsym} {self.order}\n"\
|
|
55
|
+
+ f"Degrees of freedom\n{repr(self.dofs.tolist())}\n"
|
|
56
|
+
|
|
57
|
+
def _tree_flatten(self):
|
|
58
|
+
children = (self._dofs,) # arrays / dynamic values
|
|
59
|
+
aux_data = {"n_segments": self._n_segments, "nfp": self._nfp, "stellsym": self._stellsym} # static values
|
|
60
|
+
return (children, aux_data)
|
|
61
|
+
|
|
62
|
+
@classmethod
|
|
63
|
+
def _tree_unflatten(cls, aux_data, children):
|
|
64
|
+
return cls(*children, **aux_data)
|
|
65
|
+
|
|
66
|
+
partial(jit, static_argnames=['self'])
|
|
67
|
+
def _set_gamma(self):
|
|
68
|
+
def fori_createdata(order_index: int, data: jnp.ndarray) -> jnp.ndarray:
|
|
69
|
+
return data[0] + jnp.einsum("ij,k->ikj", self._curves[:, :, 2 * order_index - 1], jnp.sin(2 * jnp.pi * order_index * self.quadpoints)) + jnp.einsum("ij,k->ikj", self._curves[:, :, 2 * order_index], jnp.cos(2 * jnp.pi * order_index * self.quadpoints)), \
|
|
70
|
+
data[1] + jnp.einsum("ij,k->ikj", self._curves[:, :, 2 * order_index - 1], 2*jnp.pi *order_index *jnp.cos(2 * jnp.pi * order_index * self.quadpoints)) + jnp.einsum("ij,k->ikj", self._curves[:, :, 2 * order_index], -2*jnp.pi *order_index *jnp.sin(2 * jnp.pi * order_index * self.quadpoints)), \
|
|
71
|
+
data[2] + jnp.einsum("ij,k->ikj", self._curves[:, :, 2 * order_index - 1], -4*jnp.pi**2*order_index**2*jnp.sin(2 * jnp.pi * order_index * self.quadpoints)) + jnp.einsum("ij,k->ikj", self._curves[:, :, 2 * order_index], -4*jnp.pi**2*order_index**2*jnp.cos(2 * jnp.pi * order_index * self.quadpoints))
|
|
72
|
+
gamma = jnp.einsum("ij,k->ikj", self._curves[:, :, 0], jnp.ones(self.n_segments))
|
|
73
|
+
gamma_dash = jnp.zeros((jnp.size(self._curves, 0), self.n_segments, 3))
|
|
74
|
+
gamma_dashdash = jnp.zeros((jnp.size(self._curves, 0), self.n_segments, 3))
|
|
75
|
+
gamma, gamma_dash, gamma_dashdash = fori_loop(1, self._order+1, fori_createdata, (gamma, gamma_dash, gamma_dashdash))
|
|
76
|
+
length = jnp.array([jnp.mean(jnp.linalg.norm(d1gamma, axis=1)) for d1gamma in gamma_dash])
|
|
77
|
+
curvature = vmap(compute_curvature)(gamma_dash, gamma_dashdash)
|
|
78
|
+
self._gamma = gamma
|
|
79
|
+
self._gamma_dash = gamma_dash
|
|
80
|
+
self._gamma_dashdash = gamma_dashdash
|
|
81
|
+
self._curvature = curvature
|
|
82
|
+
self._length = length
|
|
83
|
+
|
|
84
|
+
@property
|
|
85
|
+
def dofs(self):
|
|
86
|
+
return self._dofs
|
|
87
|
+
|
|
88
|
+
@dofs.setter
|
|
89
|
+
def dofs(self, new_dofs):
|
|
90
|
+
assert isinstance(new_dofs, jnp.ndarray)
|
|
91
|
+
assert new_dofs.ndim == 3
|
|
92
|
+
assert jnp.size(new_dofs, 1) == 3
|
|
93
|
+
assert jnp.size(new_dofs, 2) % 2 == 1
|
|
94
|
+
self._dofs = new_dofs
|
|
95
|
+
self._order = jnp.size(new_dofs, 2)//2
|
|
96
|
+
self._curves = apply_symmetries_to_curves(self.dofs, self.nfp, self.stellsym)
|
|
97
|
+
self._set_gamma()
|
|
98
|
+
|
|
99
|
+
@property
|
|
100
|
+
def curves(self):
|
|
101
|
+
return self._curves
|
|
102
|
+
|
|
103
|
+
@property
|
|
104
|
+
def order(self):
|
|
105
|
+
return self._order
|
|
106
|
+
|
|
107
|
+
@order.setter
|
|
108
|
+
def order(self, new_order):
|
|
109
|
+
assert isinstance(new_order, int)
|
|
110
|
+
assert new_order > 0
|
|
111
|
+
self._dofs = jnp.pad(self.dofs, ((0, 0), (0, 0), (0, 2*(new_order-self._order)))) if new_order > self._order else self.dofs[:, :, :2*(new_order)+1]
|
|
112
|
+
self._order = new_order
|
|
113
|
+
self._curves = apply_symmetries_to_curves(self.dofs, self.nfp, self.stellsym)
|
|
114
|
+
self._set_gamma()
|
|
115
|
+
|
|
116
|
+
@property
|
|
117
|
+
def n_segments(self):
|
|
118
|
+
return self._n_segments
|
|
119
|
+
|
|
120
|
+
@n_segments.setter
|
|
121
|
+
def n_segments(self, new_n_segments):
|
|
122
|
+
assert isinstance(new_n_segments, int)
|
|
123
|
+
assert new_n_segments > 2
|
|
124
|
+
self._n_segments = new_n_segments
|
|
125
|
+
self.quadpoints = jnp.linspace(0, 1, self._n_segments, endpoint=False)
|
|
126
|
+
self._set_gamma()
|
|
127
|
+
|
|
128
|
+
@property
|
|
129
|
+
def nfp(self):
|
|
130
|
+
return self._nfp
|
|
131
|
+
|
|
132
|
+
@nfp.setter
|
|
133
|
+
def nfp(self, new_nfp):
|
|
134
|
+
assert isinstance(new_nfp, int)
|
|
135
|
+
assert new_nfp > 0
|
|
136
|
+
self._nfp = new_nfp
|
|
137
|
+
self._curves = apply_symmetries_to_curves(self.dofs, self.nfp, self.stellsym)
|
|
138
|
+
self._set_gamma()
|
|
139
|
+
|
|
140
|
+
@property
|
|
141
|
+
def stellsym(self):
|
|
142
|
+
return self._stellsym
|
|
143
|
+
|
|
144
|
+
@stellsym.setter
|
|
145
|
+
def stellsym(self, new_stellsym):
|
|
146
|
+
assert isinstance(new_stellsym, bool)
|
|
147
|
+
self._stellsym = new_stellsym
|
|
148
|
+
self._curves = apply_symmetries_to_curves(self.dofs, self.nfp, self.stellsym)
|
|
149
|
+
self._set_gamma()
|
|
150
|
+
|
|
151
|
+
@property
|
|
152
|
+
def gamma(self):
|
|
153
|
+
return self._gamma
|
|
154
|
+
|
|
155
|
+
@property
|
|
156
|
+
def gamma_dash(self):
|
|
157
|
+
return self._gamma_dash
|
|
158
|
+
|
|
159
|
+
@property
|
|
160
|
+
def gamma_dashdash(self):
|
|
161
|
+
return self._gamma_dashdash
|
|
162
|
+
|
|
163
|
+
@property
|
|
164
|
+
def length(self):
|
|
165
|
+
return self._length
|
|
166
|
+
|
|
167
|
+
@property
|
|
168
|
+
def curvature(self):
|
|
169
|
+
return self._curvature
|
|
170
|
+
|
|
171
|
+
def save_curves(self, filename: str):
|
|
172
|
+
"""
|
|
173
|
+
Save the curves to a file
|
|
174
|
+
"""
|
|
175
|
+
with open(filename, "a") as file:
|
|
176
|
+
file.write(f"nfp stellsym order\n")
|
|
177
|
+
file.write(f"{self.nfp} {self.stellsym} {self.order}\n")
|
|
178
|
+
file.write(f"Degrees of freedom\n")
|
|
179
|
+
file.write(f"{repr(self.dofs.tolist())}\n")
|
|
180
|
+
|
|
181
|
+
def to_simsopt(self):
|
|
182
|
+
from simsopt.geo import CurveXYZFourier
|
|
183
|
+
from simsopt.field import coils_via_symmetries, Current as Current_SIMSOPT
|
|
184
|
+
|
|
185
|
+
cuves_simsopt = []
|
|
186
|
+
currents_simsopt = []
|
|
187
|
+
for dofs in self.dofs:
|
|
188
|
+
curve = CurveXYZFourier(self.n_segments, self.order)
|
|
189
|
+
curve.x = jnp.reshape(dofs, (curve.x.shape))
|
|
190
|
+
cuves_simsopt.append(curve)
|
|
191
|
+
currents_simsopt.append(Current_SIMSOPT(1))
|
|
192
|
+
coils = coils_via_symmetries(cuves_simsopt, currents_simsopt, self.nfp, self.stellsym)
|
|
193
|
+
return [c.curve for c in coils]
|
|
194
|
+
|
|
195
|
+
def plot(self, ax=None, show=True, plot_derivative=False, close=False, axis_equal=True, **kwargs):
|
|
196
|
+
def rep(data):
|
|
197
|
+
if close:
|
|
198
|
+
return jnp.concatenate((data, [data[0]]))
|
|
199
|
+
else:
|
|
200
|
+
return data
|
|
201
|
+
import matplotlib.pyplot as plt
|
|
202
|
+
if ax is None or ax.name != "3d":
|
|
203
|
+
fig = plt.figure()
|
|
204
|
+
ax = fig.add_subplot(projection='3d')
|
|
205
|
+
for gamma, gammadash in zip(self.gamma, self.gamma_dash):
|
|
206
|
+
x = rep(gamma[:, 0])
|
|
207
|
+
y = rep(gamma[:, 1])
|
|
208
|
+
z = rep(gamma[:, 2])
|
|
209
|
+
if plot_derivative:
|
|
210
|
+
xt = rep(gammadash[:, 0])
|
|
211
|
+
yt = rep(gammadash[:, 1])
|
|
212
|
+
zt = rep(gammadash[:, 2])
|
|
213
|
+
ax.plot(x, y, z, **kwargs, color='brown', linewidth=3)
|
|
214
|
+
if plot_derivative:
|
|
215
|
+
ax.quiver(x, y, z, 0.1 * xt, 0.1 * yt, 0.1 * zt, arrow_length_ratio=0.1, color="r")
|
|
216
|
+
if axis_equal:
|
|
217
|
+
fix_matplotlib_3d(ax)
|
|
218
|
+
if show:
|
|
219
|
+
plt.show()
|
|
220
|
+
|
|
221
|
+
def to_vtk(self, filename: str, close: bool = True, extra_data=None):
|
|
222
|
+
try: import numpy as np
|
|
223
|
+
except ImportError: raise ImportError("The 'numpy' library is required. Please install it using 'pip install numpy'.")
|
|
224
|
+
try: from pyevtk.hl import polyLinesToVTK
|
|
225
|
+
except ImportError: raise ImportError("The 'pyevtk' library is required. Please install it using 'pip install pyevtk'.")
|
|
226
|
+
def wrap(data):
|
|
227
|
+
return jnp.concatenate([data, jnp.array([data[0]])])
|
|
228
|
+
gammas = self.gamma
|
|
229
|
+
if close:
|
|
230
|
+
x = jnp.concatenate([wrap(gamma[:, 0]) for gamma in gammas])
|
|
231
|
+
y = jnp.concatenate([wrap(gamma[:, 1]) for gamma in gammas])
|
|
232
|
+
z = jnp.concatenate([wrap(gamma[:, 2]) for gamma in gammas])
|
|
233
|
+
ppl = jnp.asarray([gamma.shape[0]+1 for gamma in gammas])
|
|
234
|
+
else:
|
|
235
|
+
x = jnp.concatenate([gamma[:, 0] for gamma in gammas])
|
|
236
|
+
y = jnp.concatenate([gamma[:, 1] for gamma in gammas])
|
|
237
|
+
z = jnp.concatenate([gamma[:, 2] for gamma in gammas])
|
|
238
|
+
ppl = jnp.asarray([gamma.shape[0] for gamma in gammas])
|
|
239
|
+
data = jnp.concatenate([i*jnp.ones((ppl[i], )) for i in range(len(gammas))])
|
|
240
|
+
pointData = {'idx': np.array(data)}
|
|
241
|
+
if extra_data is not None:
|
|
242
|
+
pointData = {**pointData, **extra_data}
|
|
243
|
+
polyLinesToVTK(str(filename), np.array(x), np.array(y), np.array(z), pointsPerLine=np.array(ppl), pointData=pointData)
|
|
244
|
+
|
|
245
|
+
class Curves_from_simsopt(Curves):
|
|
246
|
+
# This assumes curves have all nfp and stellsym symmetries
|
|
247
|
+
def __init__(self, simsopt_curves, nfp=1, stellsym=True):
|
|
248
|
+
if isinstance(simsopt_curves, str):
|
|
249
|
+
from simsopt import load
|
|
250
|
+
bs = load(simsopt_curves)
|
|
251
|
+
simsopt_coils = bs.coils
|
|
252
|
+
simsopt_curves = [c.curve for c in simsopt_coils]
|
|
253
|
+
simsopt_curves = simsopt_curves[0:int(len(simsopt_curves)/nfp/(1+stellsym))]
|
|
254
|
+
dofs = jnp.reshape(jnp.array(
|
|
255
|
+
[curve.x for curve in simsopt_curves]
|
|
256
|
+
), (len(simsopt_curves), 3, 2*simsopt_curves[0].order+1))
|
|
257
|
+
n_segments = len(simsopt_curves[0].quadpoints)
|
|
258
|
+
super().__init__(dofs, n_segments, nfp, stellsym)
|
|
259
|
+
|
|
260
|
+
tree_util.register_pytree_node(Curves,
|
|
261
|
+
Curves._tree_flatten,
|
|
262
|
+
Curves._tree_unflatten)
|
|
263
|
+
|
|
264
|
+
class Coils(Curves):
|
|
265
|
+
def __init__(self, curves: Curves, currents: jnp.ndarray):
|
|
266
|
+
assert isinstance(curves, Curves)
|
|
267
|
+
currents = jnp.array(currents)
|
|
268
|
+
assert jnp.size(currents) == jnp.size(curves.dofs, 0)
|
|
269
|
+
super().__init__(curves.dofs, curves.n_segments, curves.nfp, curves.stellsym)
|
|
270
|
+
self.currents_scale = jnp.mean(jnp.abs(currents))
|
|
271
|
+
self._dofs_currents = currents/self.currents_scale
|
|
272
|
+
self._currents = apply_symmetries_to_currents(self._dofs_currents*self.currents_scale, self.nfp, self.stellsym)
|
|
273
|
+
|
|
274
|
+
def __str__(self):
|
|
275
|
+
return f"nfp stellsym order\n{self.nfp} {self.stellsym} {self.order}\n"\
|
|
276
|
+
+ f"Degrees of freedom\n{repr(self.dofs.tolist())}\n" \
|
|
277
|
+
+ f"Currents degrees of freedom\n{repr(self.dofs_currents.tolist())}\n" \
|
|
278
|
+
+ f"Currents scaling factor\n{self.currents_scale}\n"
|
|
279
|
+
|
|
280
|
+
def __repr__(self):
|
|
281
|
+
return f"nfp stellsym order\n{self.nfp} {self.stellsym} {self.order}\n"\
|
|
282
|
+
+ f"Degrees of freedom\n{repr(self.dofs.tolist())}\n" \
|
|
283
|
+
+ f"Currents degrees of freedom\n{repr(self.dofs_currents.tolist())}\n" \
|
|
284
|
+
+ f"Currents scaling factor\n{self.currents_scale}\n"
|
|
285
|
+
|
|
286
|
+
@property
|
|
287
|
+
def dofs_curves(self):
|
|
288
|
+
return self._dofs
|
|
289
|
+
|
|
290
|
+
@dofs_curves.setter
|
|
291
|
+
def dofs_curves(self, new_dofs_curves):
|
|
292
|
+
self.dofs = new_dofs_curves
|
|
293
|
+
|
|
294
|
+
@property
|
|
295
|
+
def dofs_currents(self):
|
|
296
|
+
return self._dofs_currents
|
|
297
|
+
|
|
298
|
+
@dofs_currents.setter
|
|
299
|
+
def dofs_currents(self, new_dofs_currents):
|
|
300
|
+
self._dofs_currents = new_dofs_currents
|
|
301
|
+
self._currents = apply_symmetries_to_currents(self._dofs_currents*self.currents_scale, self.nfp, self.stellsym)
|
|
302
|
+
|
|
303
|
+
@property
|
|
304
|
+
def x(self):
|
|
305
|
+
dofs_curves = jnp.ravel(self.dofs_curves)
|
|
306
|
+
dofs_currents = jnp.ravel(self.dofs_currents)
|
|
307
|
+
return jnp.concatenate((dofs_curves, dofs_currents))
|
|
308
|
+
|
|
309
|
+
@x.setter
|
|
310
|
+
def x(self, new_dofs):
|
|
311
|
+
old_dofs_curves = jnp.ravel(self.dofs)
|
|
312
|
+
old_dofs_currents = jnp.ravel(self.dofs_currents)
|
|
313
|
+
new_dofs_curves = new_dofs[:old_dofs_curves.shape[0]]
|
|
314
|
+
new_dofs_currents = new_dofs[old_dofs_curves.shape[0]:]
|
|
315
|
+
self.dofs_curves = jnp.reshape(new_dofs_curves, (self.dofs_curves.shape))
|
|
316
|
+
self.dofs_currents = new_dofs_currents
|
|
317
|
+
|
|
318
|
+
@property
|
|
319
|
+
def currents(self):
|
|
320
|
+
return self._currents
|
|
321
|
+
|
|
322
|
+
def _tree_flatten(self):
|
|
323
|
+
children = (Curves(self.dofs, self.n_segments, self.nfp, self.stellsym), self._dofs_currents) # arrays / dynamic values
|
|
324
|
+
aux_data = {} # static values
|
|
325
|
+
return (children, aux_data)
|
|
326
|
+
|
|
327
|
+
def save_coils(self, filename: str, text=""):
|
|
328
|
+
"""
|
|
329
|
+
Save the coils to a file
|
|
330
|
+
"""
|
|
331
|
+
with open(filename, "a") as file:
|
|
332
|
+
file.write(f"nfp stellsym order\n")
|
|
333
|
+
file.write(f"{self.nfp} {self.stellsym} {self.order}\n")
|
|
334
|
+
file.write(f"Degrees of freedom\n")
|
|
335
|
+
file.write(f"{repr(self.dofs.tolist())}\n")
|
|
336
|
+
file.write(f"Currents degrees of freedom\n")
|
|
337
|
+
file.write(f"{repr(self._dofs_currents.tolist())}\n")
|
|
338
|
+
file.write(f"Currents scaling factor\n")
|
|
339
|
+
file.write(f"{self.currents_scale}\n")
|
|
340
|
+
file.write(f"{text}\n")
|
|
341
|
+
|
|
342
|
+
def to_simsopt(self):
|
|
343
|
+
from simsopt.field import Current as Current_SIMSOPT, coils_via_symmetries
|
|
344
|
+
from simsopt.geo import CurveXYZFourier
|
|
345
|
+
cuves_simsopt = []
|
|
346
|
+
currents_simsopt = []
|
|
347
|
+
for dofs, current in zip(self.dofs_curves, self.dofs_currents*self.currents_scale):
|
|
348
|
+
curve = CurveXYZFourier(self.n_segments, self.order)
|
|
349
|
+
curve.x = jnp.reshape(dofs, (curve.x.shape))
|
|
350
|
+
cuves_simsopt.append(curve)
|
|
351
|
+
currents_simsopt.append(Current_SIMSOPT(current))
|
|
352
|
+
return coils_via_symmetries(cuves_simsopt, currents_simsopt, self.nfp, self.stellsym)
|
|
353
|
+
|
|
354
|
+
def to_json(self, filename: str):
|
|
355
|
+
data = {
|
|
356
|
+
"nfp": self.nfp,
|
|
357
|
+
"stellsym": self.stellsym,
|
|
358
|
+
"order": self.order,
|
|
359
|
+
"n_segments": self.n_segments,
|
|
360
|
+
"dofs_curves": self.dofs_curves.tolist(),
|
|
361
|
+
"dofs_currents": self.dofs_currents.tolist(),
|
|
362
|
+
}
|
|
363
|
+
import json
|
|
364
|
+
with open(filename, "w") as file:
|
|
365
|
+
json.dump(data, file)
|
|
366
|
+
|
|
367
|
+
class Coils_from_json(Coils):
|
|
368
|
+
def __init__(self, filename: str):
|
|
369
|
+
import json
|
|
370
|
+
with open(filename , "r") as file:
|
|
371
|
+
data = json.load(file)
|
|
372
|
+
super().__init__(Curves(jnp.array(data["dofs_curves"]), data["n_segments"], data["nfp"], data["stellsym"]), data["dofs_currents"])
|
|
373
|
+
|
|
374
|
+
class Coils_from_simsopt(Coils):
|
|
375
|
+
# This assumes coils have all nfp and stellsym symmetries
|
|
376
|
+
def __init__(self, simsopt_coils, nfp=1, stellsym=True):
|
|
377
|
+
if isinstance(simsopt_coils, str):
|
|
378
|
+
from simsopt import load
|
|
379
|
+
bs = load(simsopt_coils)
|
|
380
|
+
simsopt_coils = bs.coils
|
|
381
|
+
curves = [c.curve for c in simsopt_coils]
|
|
382
|
+
currents = jnp.array([c.current.get_value() for c in simsopt_coils[0:int(len(simsopt_coils)/nfp/(1+stellsym))]])
|
|
383
|
+
super().__init__(Curves_from_simsopt(curves, nfp, stellsym), currents)
|
|
384
|
+
|
|
385
|
+
tree_util.register_pytree_node(Coils,
|
|
386
|
+
Coils._tree_flatten,
|
|
387
|
+
Coils._tree_unflatten)
|
|
388
|
+
|
|
389
|
+
|
|
390
|
+
def CreateEquallySpacedCurves(n_curves: int, order: int, R: float, r: float, n_segments: int = 100,
|
|
391
|
+
nfp: int = 1, stellsym: bool = False) -> jnp.ndarray:
|
|
392
|
+
angles = (jnp.arange(n_curves) + 0.5) * (2 * jnp.pi) / ((1 + int(stellsym)) * nfp * n_curves)
|
|
393
|
+
curves = jnp.zeros((n_curves, 3, 1 + 2 * order))
|
|
394
|
+
|
|
395
|
+
curves = curves.at[:, 0, 0].set(jnp.cos(angles) * R) # x[0]
|
|
396
|
+
curves = curves.at[:, 0, 2].set(jnp.cos(angles) * r) # x[2]
|
|
397
|
+
curves = curves.at[:, 1, 0].set(jnp.sin(angles) * R) # y[0]
|
|
398
|
+
curves = curves.at[:, 1, 2].set(jnp.sin(angles) * r) # y[2]
|
|
399
|
+
curves = curves.at[:, 2, 1].set(-r) # z[1] (constant for all)
|
|
400
|
+
return Curves(curves, n_segments=n_segments, nfp=nfp, stellsym=stellsym)
|
|
401
|
+
|
|
402
|
+
def RotatedCurve(curve, phi, flip):
|
|
403
|
+
rotmat = jnp.array(
|
|
404
|
+
[[jnp.cos(phi), -jnp.sin(phi), 0],
|
|
405
|
+
[jnp.sin(phi), jnp.cos(phi), 0],
|
|
406
|
+
[0, 0, 1]]).T
|
|
407
|
+
if flip:
|
|
408
|
+
rotmat = rotmat @ jnp.array(
|
|
409
|
+
[[1, 0, 0],
|
|
410
|
+
[0, -1, 0],
|
|
411
|
+
[0, 0, -1]])
|
|
412
|
+
return curve @ rotmat
|
|
413
|
+
|
|
414
|
+
partial(jit, static_argnames=['nfp', 'stellsym'])
|
|
415
|
+
def apply_symmetries_to_curves(base_curves, nfp, stellsym):
|
|
416
|
+
flip_list = [False, True] if stellsym else [False]
|
|
417
|
+
curves = []
|
|
418
|
+
for k in range(0, nfp):
|
|
419
|
+
for flip in flip_list:
|
|
420
|
+
for i in range(len(base_curves)):
|
|
421
|
+
if k == 0 and not flip:
|
|
422
|
+
curves.append(base_curves[i])
|
|
423
|
+
else:
|
|
424
|
+
rotcurve = RotatedCurve(base_curves[i].transpose(), 2*jnp.pi*k/nfp, flip)
|
|
425
|
+
curves.append(rotcurve.transpose())
|
|
426
|
+
return jnp.array(curves)
|
|
427
|
+
|
|
428
|
+
partial(jit, static_argnames=['nfp', 'stellsym'])
|
|
429
|
+
def apply_symmetries_to_currents(base_currents, nfp, stellsym):
|
|
430
|
+
flip_list = [False, True] if stellsym else [False]
|
|
431
|
+
currents = []
|
|
432
|
+
for k in range(0, nfp):
|
|
433
|
+
for flip in flip_list:
|
|
434
|
+
for i in range(len(base_currents)):
|
|
435
|
+
current = -base_currents[i] if flip else base_currents[i]
|
|
436
|
+
currents.append(current)
|
|
437
|
+
return jnp.array(currents)
|
essos/constants.py
ADDED
|
@@ -0,0 +1,7 @@
|
|
|
1
|
+
PROTON_MASS = 1.67262192369e-27 # kg
|
|
2
|
+
NEUTRON_MASS = 1.67492749804e-27 # kg
|
|
3
|
+
ELEMENTARY_CHARGE = 1.602176634e-19 # C
|
|
4
|
+
ONE_EV = 1.602176634e-19 # J
|
|
5
|
+
ALPHA_PARTICLE_MASS = 2*PROTON_MASS + 2*NEUTRON_MASS
|
|
6
|
+
ALPHA_PARTICLE_CHARGE = 2*ELEMENTARY_CHARGE
|
|
7
|
+
FUSION_ALPHA_PARTICLE_ENERGY = 3.52e6 * ONE_EV
|