Skip to content

Commit ff9b433

Browse files
committed
Update test suite to use XGrid.from_dataset
1 parent 66ebd83 commit ff9b433

7 files changed

Lines changed: 41 additions & 46 deletions

File tree

tests/v4/test_field.py

Lines changed: 9 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,7 @@
55
import uxarray as ux
66
import xarray as xr
77

8-
from parcels import Field, UXPiecewiseConstantFace, UXPiecewiseLinearNode, VectorField, xgcm
8+
from parcels import Field, UXPiecewiseConstantFace, UXPiecewiseLinearNode, VectorField
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
@@ -15,7 +15,7 @@
1515

1616
def test_field_init_param_types():
1717
data = datasets_structured["ds_2d_left"]
18-
grid = XGrid(xgcm.Grid(data))
18+
grid = XGrid.from_dataset(data)
1919
with pytest.raises(ValueError, match="Expected `name` to be a string"):
2020
Field(name=123, data=data["data_g"], grid=grid)
2121

@@ -32,7 +32,7 @@ def test_field_init_param_types():
3232
@pytest.mark.parametrize(
3333
"data,grid",
3434
[
35-
pytest.param(ux.UxDataArray(), XGrid(xgcm.Grid(datasets_structured["ds_2d_left"])), id="uxdata-grid"),
35+
pytest.param(ux.UxDataArray(), XGrid.from_dataset(datasets_structured["ds_2d_left"]), id="uxdata-grid"),
3636
pytest.param(
3737
xr.DataArray(),
3838
UxGrid(
@@ -57,7 +57,7 @@ def test_field_incompatible_combination(data, grid):
5757
[
5858
pytest.param(
5959
datasets_structured["ds_2d_left"]["data_g"],
60-
XGrid(xgcm.Grid(datasets_structured["ds_2d_left"])),
60+
XGrid.from_dataset(datasets_structured["ds_2d_left"]),
6161
id="ds_2d_left",
6262
), # TODO: Perhaps this test should be expanded to cover more datasets?
6363
],
@@ -80,10 +80,10 @@ def test_field_init_fail_on_float_time_dim():
8080
(users are expected to use timedelta64 or datetime).
8181
"""
8282
ds = datasets_structured["ds_2d_left"].copy()
83-
ds["time"] = np.arange(0, T_structured, dtype="float64")
83+
ds["time"] = (ds["time"].dims, np.arange(0, T_structured, dtype="float64"), ds["time"].attrs)
8484

8585
data = ds["data_g"]
86-
grid = XGrid(xgcm.Grid(ds))
86+
grid = XGrid.from_dataset(ds)
8787
with pytest.raises(
8888
ValueError,
8989
match="Error getting time interval.*. Are you sure that the time dimension on the xarray dataset is stored as timedelta, datetime or cftime datetime objects\?",
@@ -100,7 +100,7 @@ def test_field_init_fail_on_float_time_dim():
100100
[
101101
pytest.param(
102102
datasets_structured["ds_2d_left"]["data_g"],
103-
XGrid(xgcm.Grid(datasets_structured["ds_2d_left"])),
103+
XGrid.from_dataset(datasets_structured["ds_2d_left"]),
104104
id="ds_2d_left",
105105
),
106106
],
@@ -119,7 +119,7 @@ def test_vectorfield_init_different_time_intervals():
119119

120120
def test_field_invalid_interpolator():
121121
ds = datasets_structured["ds_2d_left"]
122-
grid = XGrid(xgcm.Grid(ds))
122+
grid = XGrid.from_dataset(ds)
123123

124124
def invalid_interpolator_wrong_signature(self, ti, position, tau, t, z, y, invalid):
125125
return 0.0
@@ -131,7 +131,7 @@ def invalid_interpolator_wrong_signature(self, ti, position, tau, t, z, y, inval
131131

132132
def test_vectorfield_invalid_interpolator():
133133
ds = datasets_structured["ds_2d_left"]
134-
grid = XGrid(xgcm.Grid(ds))
134+
grid = XGrid.from_dataset(ds)
135135

136136
def invalid_interpolator_wrong_signature(self, ti, position, tau, t, z, y, invalid):
137137
return 0.0

tests/v4/test_fieldset.py

Lines changed: 9 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,6 @@
55
import pytest
66
import xarray as xr
77

8-
from parcels import xgcm
98
from parcels._datasets.structured.circulation_models import (
109
datasets as datasets_circulation_models, # noqa: F401
1110
) # just making sure the import works. Will eventually be used in tests
@@ -21,7 +20,7 @@
2120
@pytest.fixture
2221
def fieldset() -> FieldSet:
2322
"""Fixture to create a FieldSet object for testing."""
24-
grid = XGrid(xgcm.Grid(ds))
23+
grid = XGrid.from_dataset(ds)
2524
U = Field("U", ds["U (A grid)"], grid, mesh_type="flat")
2625
V = Field("V", ds["V (A grid)"], grid, mesh_type="flat")
2726
UV = VectorField("UV", U, V)
@@ -55,7 +54,7 @@ def test_fieldset_add_constant_field(fieldset):
5554

5655

5756
def test_fieldset_add_field(fieldset):
58-
grid = XGrid(xgcm.Grid(ds))
57+
grid = XGrid.from_dataset(ds)
5958
field = Field("test_field", ds["U (A grid)"], grid, mesh_type="flat")
6059
fieldset.add_field(field)
6160
assert fieldset.test_field == field
@@ -68,7 +67,7 @@ def test_fieldset_add_field_wrong_type(fieldset):
6867

6968

7069
def test_fieldset_add_field_already_exists(fieldset):
71-
grid = XGrid(xgcm.Grid(ds))
70+
grid = XGrid.from_dataset(ds)
7271
field = Field("test_field", ds["U (A grid)"], grid, mesh_type="flat")
7372
fieldset.add_field(field, "test_field")
7473
with pytest.raises(ValueError, match="FieldSet already has a Field with name 'test_field'"):
@@ -89,12 +88,12 @@ def test_fieldset_gridset_multiple_grids(): ...
8988

9089

9190
def test_fieldset_time_interval():
92-
grid1 = XGrid(xgcm.Grid(ds))
91+
grid1 = XGrid.from_dataset(ds)
9392
field1 = Field("field1", ds["U (A grid)"], grid1, mesh_type="flat")
9493

9594
ds2 = ds.copy()
9695
ds2["time"] = ds2["time"] + np.timedelta64(timedelta(days=1))
97-
grid2 = XGrid(xgcm.Grid(ds2))
96+
grid2 = XGrid.from_dataset(ds2)
9897
field2 = Field("field2", ds2["U (A grid)"], grid2, mesh_type="flat")
9998

10099
fieldset = FieldSet([field1, field2])
@@ -116,14 +115,14 @@ def test_fieldset_init_incompatible_calendars():
116115
ds1 = ds.copy()
117116
ds1["time"] = xr.date_range("2000", "2001", T_structured, calendar="365_day", use_cftime=True)
118117

119-
grid = XGrid(xgcm.Grid(ds1))
118+
grid = XGrid.from_dataset(ds1)
120119
U = Field("U", ds1["U (A grid)"], grid, mesh_type="flat")
121120
V = Field("V", ds1["V (A grid)"], grid, mesh_type="flat")
122121
UV = VectorField("UV", U, V)
123122

124123
ds2 = ds.copy()
125124
ds2["time"] = xr.date_range("2000", "2001", T_structured, calendar="360_day", use_cftime=True)
126-
grid2 = XGrid(xgcm.Grid(ds2))
125+
grid2 = XGrid.from_dataset(ds2)
127126
incompatible_calendar = Field("test", ds2["data_g"], grid2, mesh_type="flat")
128127

129128
with pytest.raises(CalendarError, match="Expected field '.*' to have calendar compatible with datetime object"):
@@ -133,15 +132,15 @@ def test_fieldset_init_incompatible_calendars():
133132
def test_fieldset_add_field_incompatible_calendars(fieldset):
134133
ds_test = ds.copy()
135134
ds_test["time"] = xr.date_range("2000", "2001", T_structured, calendar="360_day", use_cftime=True)
136-
grid = XGrid(xgcm.Grid(ds_test))
135+
grid = XGrid.from_dataset(ds_test)
137136
field = Field("test_field", ds_test["data_g"], grid, mesh_type="flat")
138137

139138
with pytest.raises(CalendarError, match="Expected field '.*' to have calendar compatible with datetime object"):
140139
fieldset.add_field(field, "test_field")
141140

142141
ds_test = ds.copy()
143142
ds_test["time"] = np.linspace(0, 100, T_structured, dtype="timedelta64[s]")
144-
grid = XGrid(xgcm.Grid(ds_test))
143+
grid = XGrid.from_dataset(ds_test)
145144
field = Field("test_field", ds_test["data_g"], grid, mesh_type="flat")
146145

147146
with pytest.raises(CalendarError, match="Expected field '.*' to have calendar compatible with datetime object"):

tests/v4/test_index_search.py

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,6 @@
11
import numpy as np
22
import pytest
33

4-
from parcels import xgcm
54
from parcels._datasets.structured.generic import datasets
65
from parcels._index_search import _search_indices_curvilinear_2d
76
from parcels.field import Field
@@ -13,7 +12,7 @@
1312
@pytest.fixture
1413
def field_cone():
1514
ds = datasets["2d_left_unrolled_cone"]
16-
grid = XGrid(xgcm.Grid(ds, periodic=False))
15+
grid = XGrid.from_dataset(ds)
1716
field = Field(
1817
name="test_field",
1918
data=ds["data_g"],

tests/v4/test_kernel.py

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -6,7 +6,6 @@
66
Field,
77
FieldSet,
88
ParticleSet,
9-
xgcm,
109
)
1110
from parcels._datasets.structured.generic import datasets as datasets_structured
1211
from parcels.xgrid import XGrid
@@ -16,7 +15,7 @@
1615
@pytest.fixture
1716
def fieldset() -> FieldSet:
1817
ds = datasets_structured["ds_2d_left"]
19-
grid = XGrid(xgcm.Grid(ds))
18+
grid = XGrid.from_dataset(ds)
2019
U = Field("U", ds["U (A grid)"], grid, mesh_type="flat")
2120
V = Field("V", ds["V (A grid)"], grid, mesh_type="flat")
2221
return FieldSet([U, V])

tests/v4/test_particleset.py

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -13,7 +13,6 @@
1313
ParticleSet,
1414
ParticleSetWarning,
1515
Variable,
16-
xgcm,
1716
)
1817
from parcels._datasets.structured.generic import datasets as datasets_structured
1918
from parcels.xgrid import XGrid
@@ -23,7 +22,7 @@
2322
@pytest.fixture
2423
def fieldset() -> FieldSet:
2524
ds = datasets_structured["ds_2d_left"]
26-
grid = XGrid(xgcm.Grid(ds))
25+
grid = XGrid.from_dataset(ds)
2726
U = Field("U", ds["U (A grid)"], grid, mesh_type="flat")
2827
V = Field("V", ds["V (A grid)"], grid, mesh_type="flat")
2928
return FieldSet([U, V])

tests/v4/test_particleset_execute.py

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -10,7 +10,6 @@
1010
StatusCode,
1111
UXPiecewiseConstantFace,
1212
VectorField,
13-
xgcm,
1413
)
1514
from parcels._datasets.structured.generic import datasets as datasets_structured
1615
from parcels._datasets.unstructured.generic import datasets as datasets_unstructured
@@ -22,7 +21,7 @@
2221
@pytest.fixture
2322
def fieldset() -> FieldSet:
2423
ds = datasets_structured["ds_2d_left"]
25-
grid = XGrid(xgcm.Grid(ds))
24+
grid = XGrid.from_dataset(ds)
2625
U = Field("U", ds["U (A grid)"], grid, mesh_type="flat")
2726
V = Field("V", ds["V (A grid)"], grid, mesh_type="flat")
2827
return FieldSet([U, V])

tests/v4/test_xgrid.py

Lines changed: 19 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -5,11 +5,10 @@
55
import xarray as xr
66
from numpy.testing import assert_allclose
77

8-
from parcels import xgcm
98
from parcels._datasets.structured.generic import X, Y, Z, datasets
10-
from parcels.xgrid import XGrid, _search_1d_array
9+
from parcels.xgrid import XGrid, _drop_field_data, _search_1d_array
1110

12-
GridTestCase = namedtuple("GridTestCase", ["Grid", "attr", "expected"])
11+
GridTestCase = namedtuple("GridTestCase", ["ds", "attr", "expected"])
1312

1413
test_cases = [
1514
GridTestCase(datasets["ds_2d_left"], "lon", datasets["ds_2d_left"].XG.values),
@@ -34,25 +33,25 @@ def assert_equal(actual, expected):
3433

3534
@pytest.mark.parametrize("ds, attr, expected", test_cases)
3635
def test_xgrid_properties_ground_truth(ds, attr, expected):
37-
grid = XGrid(xgcm.Grid(ds, periodic=False))
36+
grid = XGrid.from_dataset(ds)
3837
actual = getattr(grid, attr)
3938
assert_equal(actual, expected)
4039

4140

4241
@pytest.mark.parametrize("ds", [pytest.param(ds, id=key) for key, ds in datasets.items()])
43-
def test_xgrid_init_on_generic_datasets(ds):
44-
XGrid(xgcm.Grid(ds, periodic=False))
42+
def test_xgrid_from_dataset_on_generic_datasets(ds):
43+
XGrid.from_dataset(ds)
4544

4645

4746
@pytest.mark.parametrize("ds", [datasets["ds_2d_left"]])
4847
def test_xgrid_axes(ds):
49-
grid = XGrid(xgcm.Grid(ds, periodic=False))
48+
grid = XGrid.from_dataset(ds)
5049
assert grid.axes == ["Z", "Y", "X"]
5150

5251

5352
@pytest.mark.parametrize("ds", [datasets["ds_2d_left"]])
5453
def test_xgrid_get_axis_dim(ds):
55-
grid = XGrid(xgcm.Grid(ds, periodic=False))
54+
grid = XGrid.from_dataset(ds)
5655
assert grid.get_axis_dim("Z") == Z - 1
5756
assert grid.get_axis_dim("Y") == Y - 1
5857
assert grid.get_axis_dim("X") == X - 1
@@ -72,15 +71,15 @@ def test_invalid_lon_lat():
7271
ValueError,
7372
match=".*is defined on the center of the grid, but must be defined on the F points\.",
7473
):
75-
XGrid(xgcm.Grid(ds, periodic=False))
74+
XGrid.from_dataset(ds)
7675

7776
ds = datasets["ds_2d_left"].copy()
7877
ds["lon"], _ = xr.broadcast(ds["YG"], ds["XG"])
7978
with pytest.raises(
8079
ValueError,
8180
match=".*have different dimensionalities\.",
8281
):
83-
XGrid(xgcm.Grid(ds, periodic=False))
82+
XGrid.from_dataset(ds)
8483

8584
ds = datasets["ds_2d_left"].copy()
8685
ds["lon"], ds["lat"] = xr.broadcast(ds["YG"], ds["XG"])
@@ -90,7 +89,7 @@ def test_invalid_lon_lat():
9089
ValueError,
9190
match=".*must be defined on the X and Y axes and transposed to have dimensions in order of Y, X\.",
9291
):
93-
XGrid(xgcm.Grid(ds, periodic=False))
92+
XGrid.from_dataset(ds)
9493

9594

9695
@pytest.mark.parametrize(
@@ -101,7 +100,7 @@ def test_invalid_lon_lat():
101100
],
102101
) # for key, ds in datasets.items()])
103102
def test_xgrid_search_cpoints(ds):
104-
grid = XGrid(xgcm.Grid(ds, periodic=False))
103+
grid = XGrid.from_dataset(ds)
105104
lat_array, lon_array = get_2d_fpoint_mesh(grid)
106105
lat_array, lon_array = corner_to_cell_center_points(lat_array, lon_array)
107106

@@ -148,10 +147,10 @@ def test_search_1d_array(array, x, expected_xi, expected_xsi):
148147

149148

150149
@pytest.mark.parametrize(
151-
"grid, da_name, expected",
150+
"ds, da_name, expected",
152151
[
153152
pytest.param(
154-
XGrid(xgcm.Grid(datasets["ds_2d_left"], periodic=False)),
153+
datasets["ds_2d_left"],
155154
"U (C grid)",
156155
{
157156
"XG": (np.int64(0), np.float64(0.0)),
@@ -161,7 +160,7 @@ def test_search_1d_array(array, x, expected_xi, expected_xsi):
161160
id="MITgcm indexing style U (C grid)",
162161
),
163162
pytest.param(
164-
XGrid(xgcm.Grid(datasets["ds_2d_left"], periodic=False)),
163+
datasets["ds_2d_left"],
165164
"V (C grid)",
166165
{
167166
"XC": (np.int64(-1), np.float64(0.5)),
@@ -171,7 +170,7 @@ def test_search_1d_array(array, x, expected_xi, expected_xsi):
171170
id="MITgcm indexing style V (C grid)",
172171
),
173172
pytest.param(
174-
XGrid(xgcm.Grid(datasets["ds_2d_right"], periodic=False)),
173+
datasets["ds_2d_right"],
175174
"U (C grid)",
176175
{
177176
"XG": (np.int64(0), np.float64(0.0)),
@@ -181,7 +180,7 @@ def test_search_1d_array(array, x, expected_xi, expected_xsi):
181180
id="NEMO indexing style U (C grid)",
182181
),
183182
pytest.param(
184-
XGrid(xgcm.Grid(datasets["ds_2d_right"], periodic=False)),
183+
datasets["ds_2d_right"],
185184
"V (C grid)",
186185
{
187186
"XC": (np.int64(0), np.float64(0.5)),
@@ -192,10 +191,11 @@ def test_search_1d_array(array, x, expected_xi, expected_xsi):
192191
),
193192
],
194193
)
195-
def test_xgrid_localize_zero_position(grid, da_name, expected):
194+
def test_xgrid_localize_zero_position(ds, da_name, expected):
196195
"""Test localize function using left and right datasets."""
196+
grid = XGrid.from_dataset(ds)
197+
da = ds[da_name]
197198
position = grid.search(0, 0, 0)
198-
da = grid.xgcm_grid._ds[da_name]
199199

200200
local_position = grid.localize(position, da.dims)
201201
assert local_position == expected, f"Expected {expected}, got {local_position}"

0 commit comments

Comments
 (0)