Skip to content

Commit 579c26f

Browse files
committed
Update API to take `vector_field
Also add vector_field_components to the private API
1 parent 2e005b3 commit 579c26f

3 files changed

Lines changed: 60 additions & 11 deletions

File tree

src/parcels/_core/fieldset.py

Lines changed: 20 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -10,10 +10,17 @@
1010
import xarray as xr
1111

1212
from parcels._core.field import Field, VectorField
13-
from parcels._core.model import CONSTANT_FIELD_MODELS, ModelData, StructuredModelData, UnstructuredModelData
13+
from parcels._core.model import (
14+
CONSTANT_FIELD_MODELS,
15+
ModelData,
16+
StructuredModelData,
17+
TVectorFieldMapping,
18+
UnstructuredModelData,
19+
)
1420
from parcels._core.utils.string import _assert_str_and_python_varname
1521
from parcels._core.utils.time import get_datetime_type_calendar
1622
from parcels._core.utils.time import is_compatible as datetime_is_compatible
23+
from parcels._python import _MISSING, _MissingType
1724
from parcels._typing import Mesh
1825
from parcels.interpolators import (
1926
XConstantField,
@@ -201,7 +208,12 @@ def gridset(self) -> list[BaseGrid]:
201208
return grids
202209

203210
@classmethod
204-
def from_ugrid_conventions(cls, ds: ux.UxDataset, mesh: str = "spherical"):
211+
def from_ugrid_conventions(
212+
cls,
213+
ds: ux.UxDataset,
214+
mesh: str = "spherical",
215+
vector_fields: TVectorFieldMapping | None | _MissingType = _MISSING,
216+
):
205217
"""Create a FieldSet from a Parcels compliant uxarray.UxDataset.
206218
207219
This is the primary ingestion method in Parcels for structured grid datasets.
@@ -225,12 +237,15 @@ def from_ugrid_conventions(cls, ds: ux.UxDataset, mesh: str = "spherical"):
225237
-----
226238
See https://ugrid-conventions.github.io/ugrid-conventions/ for more information on the UGRID conventions.
227239
"""
228-
model = UnstructuredModelData.from_ugrid_conventions(ds, mesh)
240+
model = UnstructuredModelData.from_ugrid_conventions(ds, mesh, vector_fields)
229241
return cls([model])
230242

231243
@classmethod
232244
def from_sgrid_conventions(
233-
cls, ds: xr.Dataset, mesh: Mesh | None = None
245+
cls,
246+
ds: xr.Dataset,
247+
mesh: Mesh | None = None,
248+
vector_fields: TVectorFieldMapping | None | _MissingType = _MISSING,
234249
): # TODO: Update mesh to be discovered from the dataset metadata
235250
"""Create a FieldSet from a dataset using SGRID convention metadata.
236251
@@ -259,7 +274,7 @@ def from_sgrid_conventions(
259274
260275
See https://sgrid.github.io/sgrid/ for more information on the SGRID conventions.
261276
"""
262-
model = StructuredModelData.from_sgrid_conventions(ds, mesh)
277+
model = StructuredModelData.from_sgrid_conventions(ds, mesh, vector_fields)
263278
return cls([model])
264279

265280

src/parcels/_core/model.py

Lines changed: 28 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
from __future__ import annotations
22

33
from abc import ABC, abstractmethod
4+
from collections.abc import Mapping, Sequence
45
from typing import Any, Self
56

67
import cf_xarray # noqa: F401
@@ -18,6 +19,7 @@
1819
assert_all_field_dims_have_axis, # noqa: F401, leave import for now until decision is made # TODO v4: Make decision
1920
)
2021
from parcels._logger import logger
22+
from parcels._python import _MissingType
2123
from parcels._typing import Mesh
2224
from parcels.convert import _ds_rename_using_standard_names
2325
from parcels.interpolators import (
@@ -32,11 +34,14 @@
3234
)
3335
from parcels.interpolators._base import ScalarInterpolator, VectorInterpolator
3436

37+
TVectorFieldMapping = Mapping[str, tuple[str, str] | tuple[str, str, str]]
38+
3539

3640
class ModelData(ABC):
3741
data: Any
3842
grid: BaseGrid
3943
field_to_interpolator: dict[str, ScalarInterpolator | VectorInterpolator]
44+
vector_field_components: TVectorFieldMapping
4045

4146
@abstractmethod
4247
def construct_fields(self) -> list[Field | VectorField]: ...
@@ -79,7 +84,7 @@ def preprocess_sgrid_model_data(ds: xr.Dataset) -> xr.Dataset:
7984

8085

8186
class StructuredModelData(ModelData):
82-
def __init__(self, data: xr.Dataset, mesh: Mesh):
87+
def __init__(self, data: xr.Dataset, mesh: Mesh, vector_field_components: TVectorFieldMapping):
8388
if not isinstance(data, xr.Dataset):
8489
raise ValueError(f"Expected `data` to be an xarray.Dataset . Got {type(data)}")
8590

@@ -88,6 +93,7 @@ def __init__(self, data: xr.Dataset, mesh: Mesh):
8893

8994
self.data = data
9095
self.grid = grid
96+
self.vector_field_components = vector_field_components
9197
self.field_to_interpolator = {}
9298
self._fields: list[Field | VectorField] | None = None
9399
self.assert_valid_model_data()
@@ -133,7 +139,9 @@ def construct_fields(self) -> list[Field | VectorField]:
133139
return list(fields.values())
134140

135141
@classmethod
136-
def from_sgrid_conventions(cls, ds: xr.Dataset, mesh: Mesh | None) -> Self:
142+
def from_sgrid_conventions(
143+
cls, ds: xr.Dataset, mesh: Mesh | None, vector_fields: TVectorFieldMapping | None | _MissingType
144+
) -> Self:
137145
ds = ds.copy()
138146
if mesh is None:
139147
mesh = _get_mesh_type_from_sgrid_dataset(ds)
@@ -160,7 +168,7 @@ def from_sgrid_conventions(cls, ds: xr.Dataset, mesh: Mesh | None) -> Self:
160168
# ds["lon"] = ds[node_dimensions[0]]
161169
# ds["lat"] = ds[node_dimensions[1]]
162170

163-
model = cls(ds, mesh=mesh)
171+
model = cls(ds, mesh=mesh, vector_field_components=vector_fields)
164172
model._fields = model.construct_fields()
165173
for f in model._fields:
166174
if isinstance(f, Field):
@@ -191,13 +199,14 @@ def from_sgrid_conventions(cls, ds: xr.Dataset, mesh: Mesh | None) -> Self:
191199
),
192200
),
193201
mesh=mesh, # type:ignore
202+
vector_fields=None,
194203
)
195204
for mesh in ["flat", "spherical"]
196205
}
197206

198207

199208
class UnstructuredModelData(ModelData):
200-
def __init__(self, data: ux.UxDataset, grid: UxGrid):
209+
def __init__(self, data: ux.UxDataset, grid: UxGrid, vector_field_components: TVectorFieldMapping):
201210
if not isinstance(data, ux.UxDataset):
202211
raise ValueError(f"Expected `data` to be an uxarray.UxDataset . Got {type(data)}")
203212

@@ -206,6 +215,7 @@ def __init__(self, data: ux.UxDataset, grid: UxGrid):
206215

207216
self.data = data
208217
self.grid = grid
218+
self.vector_field_components = vector_field_components
209219
self.field_to_interpolator = {}
210220
self._fields: list[Field | VectorField] | None = None
211221

@@ -239,7 +249,9 @@ def scalar_field_names(self) -> list[str]:
239249
return list(self.data.data_vars)
240250

241251
@classmethod
242-
def from_ugrid_conventions(cls, ds: ux.UxDataset, mesh: Mesh):
252+
def from_ugrid_conventions(
253+
cls, ds: ux.UxDataset, mesh: Mesh, vector_fields: TVectorFieldMapping | None | _MissingType
254+
):
243255
ds_dims = list(ds.dims)
244256
if not all(dim in ds_dims for dim in ["time", "zf", "zc"]):
245257
raise ValueError(
@@ -276,6 +288,17 @@ def _get_mesh_type_from_sgrid_dataset(ds_sgrid: xr.Dataset) -> Mesh:
276288
return "spherical" if _is_coordinate_in_degrees(ds_sgrid[fpoint_x]) else "flat"
277289

278290

291+
def _default_vector_field_components(data_vars: Sequence[str]) -> TVectorFieldMapping:
292+
vars = set(data_vars)
293+
ret = {}
294+
295+
if {"U", "V"}.issubset(vars):
296+
ret["UV"] = ("U", "V")
297+
if {"U", "V", "W"}.issubset(vars):
298+
ret["UVW"] = ("U", "V", "W")
299+
return ret
300+
301+
279302
def _is_coordinate_in_degrees(da: xr.DataArray) -> bool:
280303
units = da.attrs.get("units")
281304
if units is None:

tests/test_fieldset.py

Lines changed: 12 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,7 @@
88

99
from parcels import Field, ParticleFile, ParticleSet, XGrid, convert
1010
from parcels._core.fieldset import FieldSet, _datetime_to_msg
11+
from parcels._core.model import _default_vector_field_components
1112
from parcels._datasets.structured.generic import datasets as datasets_structured
1213
from parcels._datasets.structured.generic import datasets_sgrid
1314
from parcels._datasets.unstructured.generic import datasets as datasets_unstructured
@@ -126,7 +127,17 @@ def test_fieldset_vectorfield_none():
126127
assert "UV" not in fset1.fields
127128

128129

129-
def test_resolve_vector_field_components(): ...
130+
@pytest.mark.parametrize(
131+
"data_vars,expected",
132+
[
133+
(["U", "V", "land_mask"], {"UV": ("U", "V")}),
134+
(["U", "V", "W", "land_mask"], {"UV": ("U", "V"), "UVW": ("U", "V", "W")}),
135+
(["field1", "field2", "field3"], {}),
136+
],
137+
)
138+
def test_default_vector_field_components(data_vars, expected):
139+
got = _default_vector_field_components(data_vars)
140+
assert got == expected
130141

131142

132143
# TODO restructure: use adding of fieldset notation to test this

0 commit comments

Comments
 (0)