Skip to content

Commit fff19b3

Browse files
Fixing execute loop for irregular dt
and adding unit test
1 parent 6aaf85b commit fff19b3

3 files changed

Lines changed: 18 additions & 5 deletions

File tree

parcels/kernel.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -381,16 +381,16 @@ def evaluate_particle(self, p, endtime):
381381
pre_dt = p.dt
382382

383383
sign_dt = np.sign(p.dt).astype(int)
384-
if sign_dt * (p.time_nextloop - endtime) > np.timedelta64(0, "ns"):
384+
if sign_dt * (endtime - p.time_nextloop) <= np.timedelta64(0, "ns"):
385385
return p
386386

387387
# TODO implement below later again
388388
# try: # Use next_dt from AdvectionRK45 if it is set
389389
# if abs(endtime - p.time_nextloop) < abs(p.next_dt) - 1e-6:
390390
# p.next_dt = abs(endtime - p.time_nextloop) * sign_dt
391391
# except AttributeError:
392-
# if abs(endtime - p.time_nextloop) < abs(p.dt) - 1e-6:
393-
# p.dt = abs(endtime - p.time_nextloop) * sign_dt
392+
if abs(endtime - p.time_nextloop) <= abs(p.dt):
393+
p.dt = abs(endtime - p.time_nextloop) * sign_dt
394394
res = self._pyfunc(p, self._fieldset, p.time_nextloop)
395395

396396
if res is None:

parcels/particleset.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -861,8 +861,8 @@ def execute(
861861
next_output = outputdt if output_file else None
862862

863863
time = start_time
864-
while time <= end_time:
865-
next_time = time + dt
864+
while time < end_time:
865+
next_time = min(time + dt, end_time) # TODO also for time-backward
866866
res = self._kernel.execute(self, endtime=next_time, dt=dt)
867867
if res == StatusCode.StopAllExecution:
868868
return StatusCode.StopAllExecution

tests/v4/test_particleset.py

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -102,6 +102,19 @@ def test_particleset_dt_type(fieldset, dt, expectation):
102102
pset.execute(runtime=np.timedelta64(10, "s"), dt=dt, pyfunc=DoNothing)
103103

104104

105+
def test_pset_starttime_not_multiple_dt(fieldset):
106+
times = [0, 1, 2]
107+
datetimes = [fieldset.U.time[0].values + np.timedelta64(t, "s") for t in times]
108+
pset = ParticleSet(fieldset, lon=[0] * len(times), lat=[0] * len(times), pclass=Particle, time=datetimes)
109+
110+
def Addlon(particle, fieldset, time): # pragma: no cover
111+
print(f"Addlon: {time} {particle.trajectory}")
112+
particle_dlon += particle.dt / np.timedelta64(1, "s") # noqa
113+
114+
pset.execute(Addlon, dt=np.timedelta64(2, "s"), runtime=np.timedelta64(8, "s"), verbose_progress=False)
115+
assert np.allclose([p.lon_nextloop for p in pset], [8 - t for t in times])
116+
117+
105118
@pytest.mark.parametrize(
106119
"runtime, expectation",
107120
[

0 commit comments

Comments
 (0)