Skip to content

Commit 97adb3b

Browse files
Adding support for negative dt (for backward tracking)
1 parent 7c47912 commit 97adb3b

2 files changed

Lines changed: 25 additions & 16 deletions

File tree

parcels/particleset.py

Lines changed: 23 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -810,6 +810,13 @@ def execute(
810810
if output_file:
811811
output_file.metadata["parcels_kernels"] = self._kernel.name
812812

813+
if np.isnat(dt) or dt is None:
814+
dt = np.timedelta64(1, "s")
815+
self._data["dt"][:] = dt
816+
sign_dt = np.sign(dt).astype(int)
817+
if sign_dt not in [-1, 1]:
818+
raise ValueError("dt must be a positive or negative np.timedelta64 object")
819+
813820
if self.fieldset.time_interval is None:
814821
start_time = np.timedelta64(0, "s") # For the execution loop, we need a start time as a timedelta object
815822
if runtime is None:
@@ -831,18 +838,23 @@ def execute(
831838
)
832839
# Ensure that the endtime uses the same type as the start_time
833840
if isinstance(endtime, self.fieldset.time_interval.left.__class__):
834-
if endtime < self.fieldset.time_interval.left:
835-
raise ValueError("The endtime must be after the start time of the fieldset.time_interval")
836-
end_time = min(endtime, self.fieldset.time_interval.right)
841+
if sign_dt > 0:
842+
if endtime < self.fieldset.time_interval.left:
843+
raise ValueError("The endtime must be after the start time of the fieldset.time_interval")
844+
end_time = min(endtime, self.fieldset.time_interval.right)
845+
else:
846+
if endtime > self.fieldset.time_interval.right:
847+
raise ValueError(
848+
"The endtime must be before the end time of the fieldset.time_interval when dt < 0"
849+
)
850+
end_time = max(endtime, self.fieldset.time_interval.left)
837851
else:
838852
raise TypeError("The endtime must be of the same type as the fieldset.time_interval start time.")
839853
else:
840-
end_time = start_time + runtime
854+
end_time = start_time + runtime * sign_dt
841855

842856
outputdt = output_file.outputdt if output_file else None
843857

844-
self._data["dt"][:] = dt
845-
846858
# Set up pbar
847859
if output_file:
848860
logger.info(f"Output files are stored in {output_file.fname}.")
@@ -853,8 +865,11 @@ def execute(
853865
next_output = outputdt if output_file else None
854866

855867
time = start_time
856-
while time < end_time:
857-
next_time = min(time + dt, end_time) # TODO also for time-backward
868+
while sign_dt * (time - end_time) < 0:
869+
if sign_dt > 0:
870+
next_time = min(time + dt, end_time)
871+
else:
872+
next_time = max(time + dt, end_time)
858873
res = self._kernel.execute(self, endtime=next_time, dt=dt)
859874
if res == StatusCode.StopAllExecution:
860875
return StatusCode.StopAllExecution

tests/v4/test_particleset_execute.py

Lines changed: 2 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -70,21 +70,15 @@ def AddLat(particle, fieldset, time): # pragma: no cover
7070

7171
@pytest.mark.parametrize(
7272
"starttime, endtime, dt",
73-
[
74-
(0, 10, 1),
75-
(0, 10, 3),
76-
(2, 16, 3),
77-
(20, 10, -1),
78-
(20, -10, -2),
79-
],
73+
[(0, 10, 1), (0, 10, 3), (2, 16, 3), (20, 10, -1), (20, 0, -2), (5, 15, None)],
8074
)
8175
def test_execution_endtime(fieldset, starttime, endtime, dt):
8276
starttime = fieldset.time_interval.left + np.timedelta64(starttime, "s")
8377
endtime = fieldset.time_interval.left + np.timedelta64(endtime, "s")
8478
dt = np.timedelta64(dt, "s")
8579
pset = ParticleSet(fieldset, time=starttime, lon=0, lat=0)
8680
pset.execute(DoNothing, endtime=endtime, dt=dt)
87-
assert np.isclose(pset.time_nextloop.values, endtime)
81+
assert abs(pset.time_nextloop - endtime) < np.timedelta64(1, "ms")
8882

8983

9084
@pytest.mark.parametrize("verbose_progress", [True, False])

0 commit comments

Comments
 (0)