Skip to content

Commit 38bbf1e

Browse files
Fix tests by passing interp_method
1 parent 0722f06 commit 38bbf1e

9 files changed

Lines changed: 58 additions & 65 deletions

tests/test_advection.py

Lines changed: 7 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -362,8 +362,8 @@ def test_stommelgyre_fieldset(kernel, rtol, grid_type):
362362
ds = stommel_gyre_dataset(grid_type=grid_type)
363363
grid = XGrid.from_dataset(ds)
364364
vector_interp_method = None if grid_type == "A" else CGrid_Velocity
365-
U = Field("U", ds["U"], grid)
366-
V = Field("V", ds["V"], grid)
365+
U = Field("U", ds["U"], grid, interp_method=XLinear)
366+
V = Field("V", ds["V"], grid, interp_method=XLinear)
367367
P = Field("P", ds["P"], grid, interp_method=XLinear)
368368
UV = VectorField("UV", U, V, vector_interp_method=vector_interp_method)
369369
fieldset = FieldSet([U, V, P, UV])
@@ -451,8 +451,8 @@ def test_nemo_curvilinear_fieldset():
451451
)
452452
grid = XGrid(xgcm_grid, mesh="spherical")
453453

454-
U = parcels.Field("U", ds["U"], grid)
455-
V = parcels.Field("V", ds["V"], grid)
454+
U = parcels.Field("U", ds["U"], grid, interp_method=XLinear)
455+
V = parcels.Field("V", ds["V"], grid, interp_method=XLinear)
456456
U.units = parcels.GeographicPolar()
457457
V.units = parcels.GeographicPolar() # U and V need GoegraphicPolar for C-Grid interpolation to work correctly
458458
UV = parcels.VectorField("UV", U, V, vector_interp_method=CGrid_Velocity)
@@ -536,9 +536,9 @@ def test_nemo_3D_curvilinear_fieldset(kernel):
536536
)
537537
grid = XGrid(xgcm_grid, mesh="spherical")
538538

539-
U = parcels.Field("U", ds["U"], grid)
540-
V = parcels.Field("V", ds["V"], grid)
541-
W = parcels.Field("W", ds["W"], grid)
539+
U = parcels.Field("U", ds["U"], grid, interp_method=XLinear)
540+
V = parcels.Field("V", ds["V"], grid, interp_method=XLinear)
541+
W = parcels.Field("W", ds["W"], grid, interp_method=XLinear)
542542
U.units = parcels.GeographicPolar()
543543
V.units = parcels.GeographicPolar() # U and V need GoegraphicPolar for C-Grid interpolation to work correctly
544544
UV = parcels.VectorField("UV", U, V, vector_interp_method=CGrid_Velocity)

tests/test_field.py

Lines changed: 12 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -9,37 +9,37 @@
99
from parcels._datasets.structured.generic import T as T_structured
1010
from parcels._datasets.structured.generic import datasets as datasets_structured
1111
from parcels._datasets.unstructured.generic import datasets as datasets_unstructured
12-
from parcels.interpolators import UXPiecewiseConstantFace, UXPiecewiseLinearNode
12+
from parcels.interpolators import UXPiecewiseConstantFace, UXPiecewiseLinearNode, XLinear
1313

1414

1515
def test_field_init_param_types():
1616
data = datasets_structured["ds_2d_left"]
1717
grid = XGrid.from_dataset(data)
1818

1919
with pytest.raises(TypeError, match="Expected a string for variable name, got int instead."):
20-
Field(name=123, data=data["data_g"], grid=grid)
20+
Field(name=123, data=data["data_g"], grid=grid, interp_method=XLinear)
2121

2222
for name in ["a b", "123"]:
2323
with pytest.raises(
2424
ValueError,
2525
match=r"Received invalid Python variable name.*: not a valid identifier. HINT: avoid using spaces, special characters, and starting with a number.",
2626
):
27-
Field(name=name, data=data["data_g"], grid=grid)
27+
Field(name=name, data=data["data_g"], grid=grid, interp_method=XLinear)
2828

2929
with pytest.raises(
3030
ValueError,
3131
match=r"Received invalid Python variable name.*: it is a reserved keyword. HINT: avoid using the following names:.*",
3232
):
33-
Field(name="while", data=data["data_g"], grid=grid)
33+
Field(name="while", data=data["data_g"], grid=grid, interp_method=XLinear)
3434

3535
with pytest.raises(
3636
ValueError,
3737
match="Expected `data` to be a uxarray.UxDataArray or xarray.DataArray",
3838
):
39-
Field(name="test", data=123, grid=grid)
39+
Field(name="test", data=123, grid=grid, interp_method=XLinear)
4040

4141
with pytest.raises(ValueError, match="Expected `grid` to be a parcels UxGrid, or parcels XGrid"):
42-
Field(name="test", data=data["data_g"], grid=123)
42+
Field(name="test", data=data["data_g"], grid=123, interp_method=XLinear)
4343

4444

4545
@pytest.mark.parametrize(
@@ -66,6 +66,7 @@ def test_field_incompatible_combination(data, grid):
6666
name="test_field",
6767
data=data,
6868
grid=grid,
69+
interp_method=XLinear,
6970
)
7071

7172

@@ -85,6 +86,7 @@ def test_field_init_structured_grid(data, grid):
8586
name="test_field",
8687
data=data,
8788
grid=grid,
89+
interp_method=XLinear,
8890
)
8991
assert field.name == "test_field"
9092
assert field.data.equals(data)
@@ -113,6 +115,7 @@ def test_field_init_fail_on_float_time_dim():
113115
name="test_field",
114116
data=data,
115117
grid=grid,
118+
interp_method=XLinear,
116119
)
117120

118121

@@ -128,7 +131,7 @@ def test_field_init_fail_on_float_time_dim():
128131
)
129132
def test_field_time_interval(data, grid):
130133
"""Test creating a field."""
131-
field = Field(name="test_field", data=data, grid=grid)
134+
field = Field(name="test_field", data=data, grid=grid, interp_method=XLinear)
132135
assert field.time_interval.left == np.datetime64("2000-01-01")
133136
assert field.time_interval.right == np.datetime64("2001-01-01")
134137

@@ -163,8 +166,8 @@ def invalid_interpolator_wrong_signature(particle_positions, grid_positions, inv
163166
return 0.0
164167

165168
# Create component fields
166-
U = Field(name="U", data=ds["data_g"], grid=grid)
167-
V = Field(name="V", data=ds["data_g"], grid=grid)
169+
U = Field(name="U", data=ds["data_g"], grid=grid, interp_method=XLinear)
170+
V = Field(name="V", data=ds["data_g"], grid=grid, interp_method=XLinear)
168171

169172
# Test invalid interpolator with wrong signature
170173
with pytest.raises(ValueError, match=".*incorrect name.*"):

tests/test_fieldset.py

Lines changed: 17 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,7 @@
1212
from parcels._datasets.structured.generic import T as T_structured
1313
from parcels._datasets.structured.generic import datasets as datasets_structured
1414
from parcels._datasets.unstructured.generic import datasets as datasets_unstructured
15+
from parcels.interpolators import XLinear
1516
from tests import utils
1617

1718
ds = datasets_structured["ds_2d_left"]
@@ -21,8 +22,8 @@
2122
def fieldset() -> FieldSet:
2223
"""Fixture to create a FieldSet object for testing."""
2324
grid = XGrid.from_dataset(ds, mesh="flat")
24-
U = Field("U", ds["U_A_grid"], grid)
25-
V = Field("V", ds["V_A_grid"], grid)
25+
U = Field("U", ds["U_A_grid"], grid, interp_method=XLinear)
26+
V = Field("V", ds["V_A_grid"], grid, interp_method=XLinear)
2627
UV = VectorField("UV", U, V)
2728

2829
return FieldSet(
@@ -65,7 +66,7 @@ def test_fieldset_add_constant_field(fieldset):
6566

6667
def test_fieldset_add_field(fieldset):
6768
grid = XGrid.from_dataset(ds, mesh="flat")
68-
field = Field("test_field", ds["U_A_grid"], grid)
69+
field = Field("test_field", ds["U_A_grid"], grid, interp_method=XLinear)
6970
fieldset.add_field(field)
7071
assert fieldset.test_field == field
7172

@@ -78,7 +79,7 @@ def test_fieldset_add_field_wrong_type(fieldset):
7879

7980
def test_fieldset_add_field_already_exists(fieldset):
8081
grid = XGrid.from_dataset(ds, mesh="flat")
81-
field = Field("test_field", ds["U_A_grid"], grid)
82+
field = Field("test_field", ds["U_A_grid"], grid, interp_method=XLinear)
8283
fieldset.add_field(field, "test_field")
8384
with pytest.raises(ValueError, match="FieldSet already has a Field with name 'test_field'"):
8485
fieldset.add_field(field, "test_field")
@@ -96,7 +97,7 @@ def test_fieldset_gridset(fieldset):
9697

9798
def test_fieldset_no_UV(tmp_zarrfile):
9899
grid = XGrid.from_dataset(ds, mesh="flat")
99-
fieldset = FieldSet([Field("P", ds["U_A_grid"], grid)])
100+
fieldset = FieldSet([Field("P", ds["U_A_grid"], grid, interp_method=XLinear)])
100101

101102
def SampleP(particles, fieldset):
102103
particles.dlon += fieldset.P[particles]
@@ -114,7 +115,7 @@ def test_fieldset_from_structured_generic_datasets(ds):
114115
grid = XGrid.from_dataset(ds, mesh="flat")
115116
fields = []
116117
for var in ds.data_vars:
117-
fields.append(Field(var, ds[var], grid))
118+
fields.append(Field(var, ds[var], grid, interp_method=XLinear))
118119

119120
fieldset = FieldSet(fields)
120121

@@ -130,12 +131,12 @@ def test_fieldset_gridset_multiple_grids(): ...
130131

131132
def test_fieldset_time_interval():
132133
grid1 = XGrid.from_dataset(ds, mesh="flat")
133-
field1 = Field("field1", ds["U_A_grid"], grid1)
134+
field1 = Field("field1", ds["U_A_grid"], grid1, interp_method=XLinear)
134135

135136
ds2 = ds.copy()
136137
ds2["time"] = (ds2["time"].dims, ds2["time"].data + np.timedelta64(timedelta(days=1)), ds2["time"].attrs)
137138
grid2 = XGrid.from_dataset(ds2, mesh="flat")
138-
field2 = Field("field2", ds2["U_A_grid"], grid2)
139+
field2 = Field("field2", ds2["U_A_grid"], grid2, interp_method=XLinear)
139140

140141
fieldset = FieldSet([field1, field2])
141142
fieldset.add_constant_field("constant_field", 1.0)
@@ -161,8 +162,8 @@ def test_fieldset_init_incompatible_calendars():
161162
)
162163

163164
grid = XGrid.from_dataset(ds1, mesh="flat")
164-
U = Field("U", ds1["U_A_grid"], grid)
165-
V = Field("V", ds1["V_A_grid"], grid)
165+
U = Field("U", ds1["U_A_grid"], grid, interp_method=XLinear)
166+
V = Field("V", ds1["V_A_grid"], grid, interp_method=XLinear)
166167
UV = VectorField("UV", U, V)
167168

168169
ds2 = ds.copy()
@@ -172,7 +173,7 @@ def test_fieldset_init_incompatible_calendars():
172173
ds2["time"].attrs,
173174
)
174175
grid2 = XGrid.from_dataset(ds2, mesh="flat")
175-
incompatible_calendar = Field("test", ds2["data_g"], grid2)
176+
incompatible_calendar = Field("test", ds2["data_g"], grid2, interp_method=XLinear)
176177

177178
with pytest.raises(CalendarError, match="Expected field '.*' to have calendar compatible with datetime object"):
178179
FieldSet([U, V, UV, incompatible_calendar])
@@ -186,7 +187,7 @@ def test_fieldset_add_field_incompatible_calendars(fieldset):
186187
ds_test["time"].attrs,
187188
)
188189
grid = XGrid.from_dataset(ds_test, mesh="flat")
189-
field = Field("test_field", ds_test["data_g"], grid)
190+
field = Field("test_field", ds_test["data_g"], grid, interp_method=XLinear)
190191

191192
with pytest.raises(CalendarError, match="Expected field '.*' to have calendar compatible with datetime object"):
192193
fieldset.add_field(field, "test_field")
@@ -198,7 +199,7 @@ def test_fieldset_add_field_incompatible_calendars(fieldset):
198199
ds_test["time"].attrs,
199200
)
200201
grid = XGrid.from_dataset(ds_test, mesh="flat")
201-
field = Field("test_field", ds_test["data_g"], grid)
202+
field = Field("test_field", ds_test["data_g"], grid, interp_method=XLinear)
202203

203204
with pytest.raises(CalendarError, match="Expected field '.*' to have calendar compatible with datetime object"):
204205
fieldset.add_field(field, "test_field")
@@ -328,8 +329,8 @@ def test_fieldset_from_fesom2_missingUV():
328329
_ = FieldSet.from_fesom2(localds)
329330
assert "Dataset has only one of the two variables 'U' and 'V'" in str(info)
330331

331-
# Intentionally create a dataset that is missing both U and V
332-
localds = ds.rename({"U": "notU", "V": "notV"})
332+
# Intentionally create a dataset that is missing the V field
333+
localds = ds.rename({"V": "notV"})
333334
with pytest.raises(ValueError) as info:
334335
_ = FieldSet.from_fesom2(localds)
335-
assert "Dataset has neither 'U' nor 'V' in potential options " in str(info)
336+
assert "Dataset has only one of the two variables 'U' and 'V'" in str(info)

tests/test_index_search.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@
77
from parcels._core.index_search import _search_indices_curvilinear_2d
88
from parcels._datasets.structured.generic import datasets
99
from parcels._tutorial import download_example_dataset
10+
from parcels.interpolators import XLinear
1011

1112

1213
@pytest.fixture
@@ -17,6 +18,7 @@ def field_cone():
1718
name="test_field",
1819
data=ds["data_g"],
1920
grid=grid,
21+
interp_method=XLinear,
2022
)
2123
return field
2224

tests/test_interpolation.py

Lines changed: 1 addition & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -9,17 +9,13 @@
99
ParticleFile,
1010
ParticleSet,
1111
StatusCode,
12-
UxGrid,
1312
Variable,
1413
VectorField,
1514
XGrid,
1615
)
1716
from parcels._core.index_search import _search_time_index
1817
from parcels._datasets.structured.generated import simple_UV_dataset
19-
from parcels._datasets.unstructured.generic import datasets as datasets_unstructured
2018
from parcels.interpolators import (
21-
UXPiecewiseConstantFace,
22-
UXPiecewiseLinearNode,
2319
XFreeslip,
2420
XLinear,
2521
XLinearInvdistLandTracer,
@@ -56,7 +52,7 @@ def field():
5652
"y": (["y"], [0.5, 1.5, 2.5, 3.5], {"axis": "Y"}),
5753
},
5854
)
59-
return Field("U", ds["U"], XGrid.from_dataset(ds))
55+
return Field("U", ds["U"], XGrid.from_dataset(ds), interp_method=XLinear)
6056

6157

6258
@pytest.mark.parametrize(
@@ -189,20 +185,6 @@ def test_interpolation_mesh_type(mesh, npart=10):
189185
assert U.eval(time, 0, lat, 0, applyConversion=False) == 1
190186

191187

192-
def test_default_interpolator_set_correctly():
193-
ds = simple_UV_dataset()
194-
grid = XGrid.from_dataset(ds)
195-
U = Field("U", ds["U"], grid)
196-
assert U.interp_method == XLinear
197-
198-
ds = datasets_unstructured["stommel_gyre_delaunay"]
199-
grid = UxGrid(grid=ds.uxgrid, z=ds.coords["nz"])
200-
U = Field("U", ds["U"], grid)
201-
assert U.interp_method == UXPiecewiseConstantFace
202-
W = Field("W", ds["W"], grid)
203-
assert W.interp_method == UXPiecewiseLinearNode
204-
205-
206188
interp_methods = {
207189
"linear": XLinear,
208190
}

tests/test_kernel.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@
1010
XGrid,
1111
)
1212
from parcels._datasets.structured.generic import datasets as datasets_structured
13+
from parcels.interpolators import XLinear
1314
from parcels.kernels import AdvectionRK4
1415
from tests.common_kernels import MoveEast, MoveNorth
1516

@@ -18,8 +19,8 @@
1819
def fieldset() -> FieldSet:
1920
ds = datasets_structured["ds_2d_left"]
2021
grid = XGrid.from_dataset(ds, mesh="flat")
21-
U = Field("U", ds["U_A_grid"], grid)
22-
V = Field("V", ds["V_A_grid"], grid)
22+
U = Field("U", ds["U_A_grid"], grid, interp_method=XLinear)
23+
V = Field("V", ds["V_A_grid"], grid, interp_method=XLinear)
2324
return FieldSet([U, V])
2425

2526

tests/test_particlefile.py

Lines changed: 6 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -22,6 +22,7 @@
2222
from parcels._core.particle import Particle, create_particle_data, get_default_particle
2323
from parcels._core.utils.time import TimeInterval
2424
from parcels._datasets.structured.generic import datasets
25+
from parcels.interpolators import XLinear
2526
from parcels.kernels import AdvectionRK4
2627
from tests.common_kernels import DoNothing
2728

@@ -31,8 +32,8 @@ def fieldset() -> FieldSet: # TODO v4: Move into a `conftest.py` file and remov
3132
"""Fixture to create a FieldSet object for testing."""
3233
ds = datasets["ds_2d_left"]
3334
grid = XGrid.from_dataset(ds)
34-
U = Field("U", ds["U_A_grid"], grid)
35-
V = Field("V", ds["V_A_grid"], grid)
35+
U = Field("U", ds["U_A_grid"], grid, XLinear)
36+
V = Field("V", ds["V_A_grid"], grid, XLinear)
3637
UV = VectorField("UV", U, V)
3738

3839
return FieldSet(
@@ -228,7 +229,9 @@ def test_write_timebackward(fieldset, tmp_zarrfile):
228229
@pytest.mark.v4alpha
229230
def test_write_xiyi(fieldset, tmp_zarrfile):
230231
fieldset.U.data[:] = 1 # set a non-zero zonal velocity
231-
fieldset.add_field(Field(name="P", data=np.zeros((3, 20)), lon=np.linspace(0, 1, 20), lat=[-2, 0, 2]))
232+
fieldset.add_field(
233+
Field(name="P", data=np.zeros((3, 20)), lon=np.linspace(0, 1, 20), lat=[-2, 0, 2], interp_method=XLinear)
234+
)
232235
dt = np.timedelta64(3600, "s")
233236

234237
particle = get_default_particle(np.float64)

tests/test_particleset.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,7 @@
1616
XGrid,
1717
)
1818
from parcels._datasets.structured.generic import datasets as datasets_structured
19+
from parcels.interpolators import XLinear
1920
from tests.common_kernels import DoNothing
2021
from tests.utils import round_and_hash_float_array
2122

@@ -24,8 +25,8 @@
2425
def fieldset() -> FieldSet:
2526
ds = datasets_structured["ds_2d_left"]
2627
grid = XGrid.from_dataset(ds, mesh="flat")
27-
U = Field("U", ds["U_A_grid"], grid)
28-
V = Field("V", ds["V_A_grid"], grid)
28+
U = Field("U", ds["U_A_grid"], grid, interp_method=XLinear)
29+
V = Field("V", ds["V_A_grid"], grid, interp_method=XLinear)
2930
return FieldSet([U, V])
3031

3132

0 commit comments

Comments
 (0)