Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
29 changes: 15 additions & 14 deletions essos/dynamics.py
Original file line number Diff line number Diff line change
Expand Up @@ -473,7 +473,8 @@ def FieldLine(t,
class Tracing():
def __init__(self, trajectories_input=None, initial_conditions=None, times_to_trace=None,
field=None, electric_field=None,model=None, maxtime: float = 1e-7, timestep: int = 1.e-8,
rtol= 1.e-7, atol = 1e-7, particles=None, condition=None,species=None,tag_gc=1.,boundary=None,rejected_steps=None):
rtol= 1.e-7, atol = 1e-7, particles=None, condition=None,species=None,tag_gc=1.,boundary=None,rejected_steps=None,
max_steps=1_000_000, progress_meter=None):

if electric_field==None:
self.electric_field = Electric_field_zero()
Expand All @@ -485,10 +486,7 @@ def __init__(self, trajectories_input=None, initial_conditions=None, times_to_tr
else:
self.field = field

if rejected_steps==None:
self.rejected_steps=100
else:
self.rejected_steps=100
self.rejected_steps = 100 if rejected_steps is None else rejected_steps

self.model = model
self.initial_conditions = initial_conditions
Expand All @@ -501,7 +499,8 @@ def __init__(self, trajectories_input=None, initial_conditions=None, times_to_tr
self.particles = particles
self.species=species
self.tag_gc=tag_gc
self.progress_meter = TqdmProgressMeter() # NoProgressMeter() # TqdmProgressMeter()
self.max_steps = max_steps
self.progress_meter = NoProgressMeter() if progress_meter is None else progress_meter
if condition is None:
self.condition = lambda t, y, args, **kwargs: False
if isinstance(field, Vmec):
Expand All @@ -528,6 +527,8 @@ def condition_BioSavart(t, y, args, **kwargs):
xx, yy, zz, _ = y
return boundary.evaluate_xyz(jnp.array([xx,yy,zz]))#<0.
self.condition = condition_BioSavart
else:
self.condition = condition
if model == 'GuidingCenter' or model=='GuidingCenterAdaptative':
self.ODE_term = ODETerm(GuidingCenter)
self.args = (self.field, self.particles,self.electric_field)
Expand Down Expand Up @@ -693,7 +694,7 @@ def update_state(state, _):
throw=False,
# adjoint=DirectAdjoint(),
#stepsize_controller = PIDController(pcoeff=0.4, icoeff=0.3, dcoeff=0, rtol=self.tol_step_size, atol=self.tol_step_size),
max_steps=10000000000,
max_steps=self.max_steps,
event = Event(self.condition),
progress_meter=self.progress_meter,
).ys
Expand All @@ -719,7 +720,7 @@ def update_state(state, _):
throw=False,
# adjoint=DirectAdjoint(),
stepsize_controller=ClipStepSizeController(controller=PIDController(pcoeff=0.1, icoeff=0.3, dcoeff=0.0, rtol=self.rtol, atol=self.atol,dtmin=dt0,dtmax=1.e-4,force_dtmin=True),step_ts=self.times,store_rejected_steps=self.rejected_steps),
max_steps=10000000000,
max_steps=self.max_steps,
event = Event(self.condition),
progress_meter=self.progress_meter,
).ys
Expand All @@ -743,7 +744,7 @@ def update_state(state, _):
saveat=SaveAt(ts=self.times),
throw=False,
# adjoint=DirectAdjoint(),
max_steps=10000000000,
max_steps=self.max_steps,
event = Event(self.condition),
progress_meter=self.progress_meter,
).ys
Expand All @@ -767,7 +768,7 @@ def update_state(state, _):
saveat=SaveAt(ts=self.times),
throw=False,
# adjoint=DirectAdjoint(),
max_steps=10000000000,
max_steps=self.max_steps,
event = Event(self.condition),
progress_meter=self.progress_meter,
).ys
Expand All @@ -793,7 +794,7 @@ def update_state(state, _):
throw=False,
# adjoint=DirectAdjoint(),
stepsize_controller = PIDController(pcoeff=0.4, icoeff=0.3, dcoeff=0, rtol=self.tol_step_size, atol=self.tol_step_size,dtmin=dt0),
max_steps=10000000000,
max_steps=self.max_steps,
event = Event(self.condition),
progress_meter=self.progress_meter,
).ys
Expand All @@ -813,7 +814,7 @@ def update_state(state, _):
# adjoint=DirectAdjoint(),
progress_meter=self.progress_meter,
stepsize_controller = PIDController(pcoeff=0.4, icoeff=0.3, dcoeff=0, rtol=self.rtol, atol=self.atol),
max_steps=10000000000,
max_steps=self.max_steps,
event = Event(self.condition)
).ys
elif self.model == 'FieldLineAdaptative' :
Expand All @@ -832,7 +833,7 @@ def update_state(state, _):
# adjoint=DirectAdjoint(),
progress_meter=self.progress_meter,
stepsize_controller = PIDController(pcoeff=0.4, icoeff=0.3, dcoeff=0, rtol=self.rtol, atol=self.atol),
max_steps=10000000000,
max_steps=self.max_steps,
event = Event(self.condition)
).ys
#Fixed guiding center
Expand All @@ -851,7 +852,7 @@ def update_state(state, _):
throw=False,
# adjoint=DirectAdjoint(),
progress_meter=self.progress_meter,
max_steps=10000000000,
max_steps=self.max_steps,
event = Event(self.condition)
).ys
return trajectory
Expand Down
12 changes: 10 additions & 2 deletions tests/test_dynamics.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import pytest
import jax.numpy as jnp
from diffrax import NoProgressMeter
from essos.constants import ALPHA_PARTICLE_MASS, ALPHA_PARTICLE_CHARGE, FUSION_ALPHA_PARTICLE_ENERGY,ELECTRON_MASS,PROTON_MASS
from essos.dynamics import Particles, GuidingCenter, Lorentz, FieldLine, Tracing
from essos.background_species import BackgroundSpecies
Expand Down Expand Up @@ -131,15 +132,22 @@ def test_field_line(field):
assert result.shape == (3,)

def test_tracing_initialization(field, particles,electric_field):
def condition(t, y, args, **kwargs):
return False

x = jnp.linspace(1, 2, particles.nparticles)
y = jnp.zeros(particles.nparticles)
z = jnp.zeros(particles.nparticles)
initial_conditions =jnp.array([x, y, z]).T
tracing = Tracing(initial_conditions=initial_conditions, field=field,electric_field=electric_field, model='GuidingCenter', particles=particles, times_to_trace=200)
tracing = Tracing(initial_conditions=initial_conditions, field=field,electric_field=electric_field, model='GuidingCenter', particles=particles, times_to_trace=200, condition=condition, rejected_steps=7, max_steps=123)
assert tracing.field == field
assert tracing.model == 'GuidingCenter'
assert tracing.initial_conditions.shape == (particles.nparticles, 4)
assert tracing.times.shape == (200,)
assert tracing.condition is condition
assert tracing.rejected_steps == 7
assert tracing.max_steps == 123
assert isinstance(tracing.progress_meter, NoProgressMeter)

def test_tracing_trace(field, particles,electric_field):
x = jnp.linspace(1, 2, particles.nparticles)
Expand Down Expand Up @@ -221,4 +229,4 @@ def test_tracing_trace_collisions_adaptative(field, particles,electric_field):
assert trajectories.shape == (particles.nparticles, 200, 5)

if __name__ == "__main__":
pytest.main()
pytest.main()
Loading