Skip to content

Commit 180e9a5

Browse files
Adding more tests
1 parent b6fa426 commit 180e9a5

3 files changed

Lines changed: 108 additions & 1 deletion

File tree

parcels/particleset.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -729,7 +729,9 @@ def execute(
729729
if output_file:
730730
output_file.metadata["parcels_kernels"] = self._kernel.name
731731

732-
if np.isnat(dt) or dt is None:
732+
if (dt is not None) and (not isinstance(dt, np.timedelta64)):
733+
raise TypeError("dt must be a np.timedelta64 object")
734+
if dt is None or np.isnat(dt):
733735
dt = np.timedelta64(1, "s")
734736
self._data["dt"][:] = dt
735737
sign_dt = np.sign(dt).astype(int)

tests/v4/test_kernel.py

Lines changed: 75 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,75 @@
1+
import numpy as np
2+
import pytest
3+
4+
from parcels import (
5+
AdvectionRK4,
6+
Field,
7+
FieldSet,
8+
ParticleSet,
9+
xgcm,
10+
)
11+
from parcels._datasets.structured.generic import datasets as datasets_structured
12+
from parcels.xgrid import XGrid
13+
from tests.common_kernels import MoveEast, MoveNorth
14+
15+
16+
@pytest.fixture
17+
def fieldset() -> FieldSet:
18+
ds = datasets_structured["ds_2d_left"]
19+
grid = XGrid(xgcm.Grid(ds))
20+
U = Field("U", ds["U (A grid)"], grid, mesh_type="flat")
21+
V = Field("V", ds["V (A grid)"], grid, mesh_type="flat")
22+
return FieldSet([U, V])
23+
24+
25+
def test_multi_kernel_reuse_varnames(fieldset):
26+
pset = ParticleSet(fieldset, lon=[0.5], lat=[0.5])
27+
28+
# Testing for merging of two Kernels with the same variable declared
29+
def MoveEast1(particle, fieldset, time): # pragma: no cover
30+
add_lon = 0.2
31+
particle_dlon += add_lon # noqa
32+
33+
def MoveEast2(particle, fieldset, time): # pragma: no cover
34+
particle_dlon += add_lon # noqa
35+
36+
pset.execute([MoveEast1, MoveEast2], runtime=np.timedelta64(2, "s"))
37+
assert np.allclose(pset.lon, [0.9], atol=1e-5) # should be 0.5 + 0.2 + 0.2 = 0.9
38+
39+
40+
def test_combined_kernel_from_list(fieldset):
41+
"""
42+
Test pset.Kernel(List[function])
43+
44+
Tests that a Kernel can be created from a list functions, or a list of
45+
mixed functions and kernel objects.
46+
"""
47+
pset = ParticleSet(fieldset, lon=[0.5], lat=[0.5])
48+
kernels_single = pset.Kernel([AdvectionRK4])
49+
kernels_functions = pset.Kernel([AdvectionRK4, MoveEast, MoveNorth])
50+
51+
# Check if the kernels were combined correctly
52+
assert kernels_single.funcname == "AdvectionRK4"
53+
assert kernels_functions.funcname == "AdvectionRK4MoveEastMoveNorth"
54+
55+
56+
def test_combined_kernel_from_list_error_checking(fieldset):
57+
"""
58+
Test pset.Kernel(List[function])
59+
60+
Tests that various error cases raise appropriate messages.
61+
"""
62+
pset = ParticleSet(fieldset, lon=[0.5], lat=[0.5])
63+
64+
# Test that list has to be non-empty
65+
with pytest.raises(ValueError):
66+
pset.Kernel([])
67+
68+
# Test that list has to be all functions
69+
with pytest.raises(ValueError):
70+
pset.Kernel([AdvectionRK4, "something else"])
71+
72+
# Can't mix kernel objects and functions in list
73+
with pytest.raises(ValueError):
74+
kernels_mixed = pset.Kernel([pset.Kernel(AdvectionRK4), MoveEast, MoveNorth])
75+
assert kernels_mixed.funcname == "AdvectionRK4MoveEastMoveNorth"

tests/v4/test_particleset_execute.py

Lines changed: 30 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -81,6 +81,36 @@ def test_execution_endtime(fieldset, starttime, endtime, dt):
8181
assert abs(pset.time_nextloop - endtime) < np.timedelta64(1, "ms")
8282

8383

84+
@pytest.mark.parametrize(
85+
"starttime, runtime, dt",
86+
[(0, 10, 1), (0, 10, 3), (2, 16, 3), (20, 10, -1), (20, 0, -2), (5, 15, None)],
87+
)
88+
def test_execution_runtime(fieldset, starttime, runtime, dt):
89+
starttime = fieldset.time_interval.left + np.timedelta64(starttime, "s")
90+
runtime = np.timedelta64(runtime, "s")
91+
sign_dt = 1 if dt is None else np.sign(dt)
92+
dt = np.timedelta64(dt, "s")
93+
pset = ParticleSet(fieldset, time=starttime, lon=0, lat=0)
94+
pset.execute(DoNothing, runtime=runtime, dt=dt)
95+
assert abs(pset.time_nextloop - starttime - runtime * sign_dt) < np.timedelta64(1, "ms")
96+
97+
98+
def test_execution_fail_python_exception(fieldset, npart=10):
99+
pset = ParticleSet(fieldset, lon=np.linspace(0, 1, npart), lat=np.linspace(1, 0, npart))
100+
101+
def PythonFail(particle, fieldset, time): # pragma: no cover
102+
if particle.time >= fieldset.time_interval.left + np.timedelta64(10, "s"):
103+
raise RuntimeError("Enough is enough!")
104+
else:
105+
pass
106+
107+
with pytest.raises(RuntimeError):
108+
pset.execute(PythonFail, runtime=np.timedelta64(20, "s"), dt=np.timedelta64(2, "s"))
109+
assert len(pset) == npart
110+
assert pset.time[0] == fieldset.time_interval.left + np.timedelta64(10, "s")
111+
assert all([time == fieldset.time_interval.left + np.timedelta64(8, "s") for time in pset.time[1:]])
112+
113+
84114
@pytest.mark.parametrize("verbose_progress", [True, False])
85115
def test_uxstommelgyre_pset_execute(verbose_progress):
86116
ds = datasets_unstructured["stommel_gyre_delaunay"]

0 commit comments

Comments
 (0)