Skip to content

Commit 5c1b8cb

Browse files
VeckoTheGeckopre-commit-ci[bot]erikvansebille
authored
Add vector_fields parameter to FieldSet creation methods (#2715)
* Move tests and rename * Update field check * Remove test stub No longer relevant * Add tests for custom vectorfields * Remove default mesh on ModelData from_{s,u}grid_conventions These methods arent public API and these defaults are set on the FieldSet class * Add MISSING sentinel value * Update API to take `vector_field Also add vector_field_components to the private API * Remove duplicate function * Update StructuredModelData.construct_fields() * Update naming * Update unstructured code to work with custom vectorfields * Update docstring * Rename sentinel value * Fix tests * Fix mypy issues * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Move TVectorField to _typing module * Improve validation of vector_fields * Remove None as option for vector_fields * Improve test_fieldset_add Now also has vectorfields * Update src/parcels/_core/fieldset.py Co-authored-by: Erik van Sebille <e.vansebille@uu.nl> * Update src/parcels/_core/fieldset.py Co-authored-by: Erik van Sebille <e.vansebille@uu.nl> --------- Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com> Co-authored-by: Erik van Sebille <e.vansebille@uu.nl>
1 parent c2f31e4 commit 5c1b8cb

6 files changed

Lines changed: 260 additions & 75 deletions

File tree

src/parcels/_core/fieldset.py

Lines changed: 29 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -9,12 +9,18 @@
99
import uxarray as ux
1010
import xarray as xr
1111

12+
import parcels._typing as ptyping
1213
from parcels._core.field import Field, VectorField
13-
from parcels._core.model import CONSTANT_FIELD_MODELS, ModelData, StructuredModelData, UnstructuredModelData
14+
from parcels._core.model import (
15+
CONSTANT_FIELD_MODELS,
16+
ModelData,
17+
StructuredModelData,
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
17-
from parcels._typing import Mesh
23+
from parcels._python import NOTSET, NotSetType
1824
from parcels.interpolators import (
1925
XConstantField,
2026
)
@@ -144,7 +150,7 @@ def add_field(self, field: Field, name: str | None = None):
144150

145151
self.fields[name] = field
146152

147-
def add_constant_field(self, name: str, value, mesh: Mesh = "spherical"):
153+
def add_constant_field(self, name: str, value, mesh: ptyping.Mesh = "spherical"):
148154
"""Wrapper function to add a Field that is constant in space,
149155
useful e.g. when using constant horizontal diffusivity
150156
@@ -201,7 +207,12 @@ def gridset(self) -> list[BaseGrid]:
201207
return grids
202208

203209
@classmethod
204-
def from_ugrid_conventions(cls, ds: ux.UxDataset, mesh: str = "spherical"):
210+
def from_ugrid_conventions(
211+
cls,
212+
ds: ux.UxDataset,
213+
mesh: str = "spherical",
214+
vector_fields: ptyping.VectorFields | NotSetType = NOTSET,
215+
):
205216
"""Create a FieldSet from a Parcels compliant uxarray.UxDataset.
206217
207218
This is the primary ingestion method in Parcels for structured grid datasets.
@@ -215,6 +226,10 @@ def from_ugrid_conventions(cls, ds: ux.UxDataset, mesh: str = "spherical"):
215226
----------
216227
ds : uxarray.UxDataset
217228
uxarray.UxDataset as obtained from the uxarray package but with appropriate named vertical dimensions
229+
vector_fields : Mapping[str, tuple[str, ...]], optional
230+
Mapping of vector field names to tuples of component variable names in the dataset.
231+
For example, ``{"UV": ("U", "V"), "UVW": ("U", "V", "W")}``.
232+
If omitted (default), vector fields are auto-discovered from standard variable names (``U``/``V``/``W``).
218233
219234
Returns
220235
-------
@@ -225,12 +240,15 @@ def from_ugrid_conventions(cls, ds: ux.UxDataset, mesh: str = "spherical"):
225240
-----
226241
See https://ugrid-conventions.github.io/ugrid-conventions/ for more information on the UGRID conventions.
227242
"""
228-
model = UnstructuredModelData.from_ugrid_conventions(ds, mesh)
243+
model = UnstructuredModelData.from_ugrid_conventions(ds, mesh, vector_fields)
229244
return cls([model])
230245

231246
@classmethod
232247
def from_sgrid_conventions(
233-
cls, ds: xr.Dataset, mesh: Mesh | None = None
248+
cls,
249+
ds: xr.Dataset,
250+
mesh: ptyping.Mesh | None = None,
251+
vector_fields: ptyping.VectorFields | NotSetType = NOTSET,
234252
): # TODO: Update mesh to be discovered from the dataset metadata
235253
"""Create a FieldSet from a dataset using SGRID convention metadata.
236254
@@ -245,6 +263,10 @@ def from_sgrid_conventions(
245263
mesh : str
246264
String indicating the type of mesh coordinates used during
247265
velocity interpolation. Options are "spherical" or "flat".
266+
vector_fields : Mapping[str, tuple[str, ...]], optional
267+
Mapping of vector field names to tuples of component variable names in the dataset.
268+
For example, ``{"UV": ("U", "V"), "UVW": ("U", "V", "W")}``.
269+
If omitted (default), vector fields are auto-discovered from standard variable names (``U``/``V``/``W``).
248270
249271
Returns
250272
-------
@@ -259,7 +281,7 @@ def from_sgrid_conventions(
259281
260282
See https://sgrid.github.io/sgrid/ for more information on the SGRID conventions.
261283
"""
262-
model = StructuredModelData.from_sgrid_conventions(ds, mesh)
284+
model = StructuredModelData.from_sgrid_conventions(ds, mesh, vector_fields)
263285
return cls([model])
264286

265287

@@ -356,9 +378,3 @@ def _format_calendar_error_message(field: Field | VectorField, reference_datetim
356378
],
357379
"W": ["upward_sea_water_velocity", "vertical_sea_water_velocity"],
358380
}
359-
360-
361-
def _is_agrid(ds: xr.Dataset) -> bool:
362-
# check if U and V are defined on the same dimensions
363-
# if yes, interpret as A grid
364-
return set(ds["U"].dims) == set(ds["V"].dims)

src/parcels/_core/model.py

Lines changed: 94 additions & 39 deletions
Original file line numberDiff line numberDiff line change
@@ -1,13 +1,15 @@
11
from __future__ import annotations
22

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

67
import cf_xarray # noqa: F401
78
import uxarray as ux
89
import xarray as xr
910

1011
import parcels._sgrid as sgrid
12+
import parcels._typing as ptyping
1113
from parcels._core.basegrid import BaseGrid
1214
from parcels._core.field import Field, VectorField
1315
from parcels._core.utils.time import TimeInterval
@@ -18,6 +20,7 @@
1820
assert_all_field_dims_have_axis, # noqa: F401, leave import for now until decision is made # TODO v4: Make decision
1921
)
2022
from parcels._logger import logger
23+
from parcels._python import NOTSET, NotSetType
2124
from parcels._typing import Mesh
2225
from parcels.convert import _ds_rename_using_standard_names
2326
from parcels.interpolators import (
@@ -37,6 +40,7 @@ class ModelData(ABC):
3740
data: Any
3841
grid: BaseGrid
3942
field_to_interpolator: dict[str, ScalarInterpolator | VectorInterpolator]
43+
vector_field_components: ptyping.VectorFields
4044

4145
@abstractmethod
4246
def construct_fields(self) -> list[Field | VectorField]: ...
@@ -79,7 +83,7 @@ def preprocess_sgrid_model_data(ds: xr.Dataset) -> xr.Dataset:
7983

8084

8185
class StructuredModelData(ModelData):
82-
def __init__(self, data: xr.Dataset, mesh: Mesh):
86+
def __init__(self, data: xr.Dataset, mesh: Mesh, vector_field_components: ptyping.VectorFields):
8387
if not isinstance(data, xr.Dataset):
8488
raise ValueError(f"Expected `data` to be an xarray.Dataset . Got {type(data)}")
8589

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

8993
self.data = data
9094
self.grid = grid
95+
self.vector_field_components = vector_field_components
9196
self.field_to_interpolator = {}
9297
self._fields: list[Field | VectorField] | None = None
9398
self.assert_valid_model_data()
@@ -110,30 +115,25 @@ def construct_fields(self) -> list[Field | VectorField]:
110115
single_fields: dict[str, Field] = {}
111116
vector_fields: dict[str, VectorField] = {}
112117
scalar_field_names = self.scalar_field_names
113-
if "U" in scalar_field_names and "V" in scalar_field_names:
114-
interp_method = XLinear_Velocity() if _is_agrid(self.data) else CGrid_Velocity()
115-
single_fields["U"] = Field("U", self)
116-
single_fields["V"] = Field("V", self)
117-
vector_fields["UV"] = VectorField("UV", single_fields["U"], single_fields["V"], interp_method=interp_method)
118-
119-
if "W" in scalar_field_names:
120-
single_fields["W"] = Field("W", self)
121-
vector_fields["UVW"] = VectorField(
122-
"UVW",
123-
single_fields["U"],
124-
single_fields["V"],
125-
single_fields["W"],
126-
interp_method=interp_method,
127-
)
128118

129-
fields: dict[str, Field | VectorField] = {**single_fields, **vector_fields}
130-
for varname in set(scalar_field_names) - set(fields.keys()):
131-
fields[varname] = Field(str(varname), self)
119+
for varname in set(scalar_field_names):
120+
single_fields[varname] = Field(str(varname), self)
121+
122+
for vfield_name, components in self.vector_field_components.items():
123+
interp_method = (
124+
XLinear_Velocity() if _is_agrid(self.data, u=components[0], v=components[1]) else CGrid_Velocity()
125+
)
126+
127+
component_fields = [single_fields[name] for name in components]
128+
vector_fields[vfield_name] = VectorField(vfield_name, *component_fields, interp_method=interp_method) # type:ignore[misc,arg-type]
132129

130+
fields: dict[str, Field | VectorField] = {**single_fields, **vector_fields}
133131
return list(fields.values())
134132

135133
@classmethod
136-
def from_sgrid_conventions(cls, ds: xr.Dataset, mesh: Mesh | None = None) -> Self:
134+
def from_sgrid_conventions(
135+
cls, ds: xr.Dataset, mesh: Mesh | None, vector_fields: ptyping.VectorFields | NotSetType
136+
) -> Self:
137137
ds = ds.copy()
138138
if mesh is None:
139139
mesh = _get_mesh_type_from_sgrid_dataset(ds)
@@ -160,14 +160,56 @@ def from_sgrid_conventions(cls, ds: xr.Dataset, mesh: Mesh | None = None) -> Sel
160160
# ds["lon"] = ds[node_dimensions[0]]
161161
# ds["lat"] = ds[node_dimensions[1]]
162162

163-
model = cls(ds, mesh=mesh)
163+
vector_fields = resolve_vector_fields(ds, vector_fields)
164+
assert_valid_vector_fields(ds, vector_fields)
165+
166+
model = cls(ds, mesh=mesh, vector_field_components=vector_fields)
164167
model._fields = model.construct_fields()
165168
for f in model._fields:
166169
if isinstance(f, Field):
167170
f.interp_method = XLinear()
168171
return model
169172

170173

174+
def resolve_vector_fields(ds: xr.Dataset, vector_fields: ptyping.VectorFields | NotSetType) -> ptyping.VectorFields:
175+
if vector_fields is NOTSET: # i.e., the default vectorfield discovery behaviour
176+
return _default_vector_field_components(list(ds.data_vars))
177+
return vector_fields
178+
179+
180+
def assert_valid_vector_fields(ds: xr.Dataset, vector_fields: ptyping.VectorFields) -> None:
181+
if not isinstance(vector_fields, dict):
182+
raise ValueError(f"vector_fields must be a dictionary. Got {type(vector_fields)=!r}.")
183+
184+
for vfield_name, components in vector_fields.items():
185+
if not isinstance(vfield_name, str):
186+
raise ValueError(
187+
f"Invalid `vector_fields` argument. Vector field name in `vector_fields` should be a string. Got field name {vfield_name!r}."
188+
)
189+
if not (2 <= len(components) <= 3):
190+
raise ValueError(
191+
f"Invalid `vector_fields` argument. Vector fields must have either 2 or 3 components. Vector field {vfield_name} has {len(components)} components."
192+
)
193+
for c in components:
194+
if not isinstance(c, str):
195+
raise ValueError(
196+
f"Invalid `vector_fields` argument. Component names must be strings. Got component name of value {c!r}."
197+
)
198+
199+
assert_vector_field_components_in_dataset(ds, vector_fields)
200+
return
201+
202+
203+
def assert_vector_field_components_in_dataset(ds: xr.Dataset, vector_fields: ptyping.VectorFields) -> None:
204+
for components in vector_fields.values():
205+
for c in components:
206+
if c not in ds.data_vars:
207+
raise ValueError(
208+
f"Field component '{c}' not present in the source dataset, but is listed in {vector_fields=!r}. This component cannot be used in this mapping."
209+
)
210+
return
211+
212+
171213
CONSTANT_FIELD_MODELS = {
172214
mesh: StructuredModelData.from_sgrid_conventions(
173215
xr.Dataset(
@@ -191,13 +233,14 @@ def from_sgrid_conventions(cls, ds: xr.Dataset, mesh: Mesh | None = None) -> Sel
191233
),
192234
),
193235
mesh=mesh, # type:ignore
236+
vector_fields={},
194237
)
195238
for mesh in ["flat", "spherical"]
196239
}
197240

198241

199242
class UnstructuredModelData(ModelData):
200-
def __init__(self, data: ux.UxDataset, grid: UxGrid):
243+
def __init__(self, data: ux.UxDataset, grid: UxGrid, vector_field_components: ptyping.VectorFields):
201244
if not isinstance(data, ux.UxDataset):
202245
raise ValueError(f"Expected `data` to be an uxarray.UxDataset . Got {type(data)}")
203246

@@ -206,28 +249,25 @@ def __init__(self, data: ux.UxDataset, grid: UxGrid):
206249

207250
self.data = data
208251
self.grid = grid
252+
self.vector_field_components = vector_field_components
209253
self.field_to_interpolator = {}
210254
self._fields: list[Field | VectorField] | None = None
211255

212256
def construct_fields(self) -> list[Field | VectorField]:
213257
single_fields: dict[str, Field] = {}
214258
vector_fields: dict[str, VectorField] = {}
215259
scalar_field_names = self.scalar_field_names
216-
if "U" in scalar_field_names and "V" in scalar_field_names:
217-
single_fields["U"] = Field("U", self)
218-
single_fields["V"] = Field("V", self)
219-
vector_fields["UV"] = VectorField("UV", single_fields["U"], single_fields["V"], interp_method=Ux_Velocity())
220-
221-
if "W" in scalar_field_names:
222-
single_fields["W"] = Field("W", self)
223-
vector_fields["UVW"] = VectorField(
224-
"UVW", single_fields["U"], single_fields["V"], single_fields["W"], interp_method=Ux_Velocity()
225-
)
226260

227-
fields: dict[str, Field | VectorField] = {**single_fields, **vector_fields}
228-
for varname in set(scalar_field_names) - set(single_fields.keys()):
229-
fields[varname] = Field(str(varname), self)
261+
for varname in set(scalar_field_names):
262+
single_fields[varname] = Field(str(varname), self)
230263

264+
for vfield_name, components in self.vector_field_components.items():
265+
interp_method = Ux_Velocity()
266+
267+
component_fields = [single_fields[name] for name in components]
268+
vector_fields[vfield_name] = VectorField(vfield_name, *component_fields, interp_method=interp_method) # type:ignore[misc, arg-type]
269+
270+
fields: dict[str, Field | VectorField] = {**single_fields, **vector_fields}
231271
return list(fields.values())
232272

233273
def assert_valid_field_data(self, field_data: ux.UxDataArray) -> None:
@@ -239,7 +279,7 @@ def scalar_field_names(self) -> list[str]:
239279
return list(self.data.data_vars)
240280

241281
@classmethod
242-
def from_ugrid_conventions(cls, ds: ux.UxDataset, mesh: str = "spherical"):
282+
def from_ugrid_conventions(cls, ds: ux.UxDataset, mesh: Mesh, vector_fields: ptyping.VectorFields | NotSetType):
243283
ds_dims = list(ds.dims)
244284
if not all(dim in ds_dims for dim in ["time", "zf", "zc"]):
245285
raise ValueError(
@@ -248,7 +288,11 @@ def from_ugrid_conventions(cls, ds: ux.UxDataset, mesh: str = "spherical"):
248288

249289
grid = UxGrid(ds.uxgrid, z=ds.coords["zf"], mesh=mesh)
250290
ds = _discover_ux_U_and_V(ds)
251-
model = cls(ds, grid)
291+
292+
vector_fields = resolve_vector_fields(ds, vector_fields)
293+
assert_valid_vector_fields(ds, vector_fields)
294+
295+
model = cls(ds, grid, vector_fields)
252296
model._fields = model.construct_fields()
253297
for f in model._fields:
254298
if isinstance(f, Field):
@@ -276,6 +320,17 @@ def _get_mesh_type_from_sgrid_dataset(ds_sgrid: xr.Dataset) -> Mesh:
276320
return "spherical" if _is_coordinate_in_degrees(ds_sgrid[fpoint_x]) else "flat"
277321

278322

323+
def _default_vector_field_components(data_vars: Sequence[Hashable]) -> ptyping.VectorFields:
324+
vars = set(data_vars)
325+
ret: ptyping.VectorFields = {}
326+
327+
if {"U", "V"}.issubset(vars):
328+
ret["UV"] = ("U", "V")
329+
if {"U", "V", "W"}.issubset(vars):
330+
ret["UVW"] = ("U", "V", "W")
331+
return ret
332+
333+
279334
def _is_coordinate_in_degrees(da: xr.DataArray) -> bool:
280335
units = da.attrs.get("units")
281336
if units is None:
@@ -366,10 +421,10 @@ def _select_uxinterpolator(da: ux.UxDataArray):
366421
return None
367422

368423

369-
def _is_agrid(ds: xr.Dataset) -> bool:
424+
def _is_agrid(ds: xr.Dataset, u: str, v: str) -> bool:
370425
# check if U and V are defined on the same dimensions
371426
# if yes, interpret as A grid
372-
return set(ds["U"].dims) == set(ds["V"].dims)
427+
return set(ds[u].dims) == set(ds[v].dims)
373428

374429

375430
def _get_time_interval(data: xr.DataArray | ux.UxDataArray) -> TimeInterval | None:

src/parcels/_python.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,11 +1,15 @@
11
# Generic Python helpers
2+
import enum
23
import inspect
34
from collections.abc import Callable, Mapping
45
from typing import TypeVar
56

67
K = TypeVar("K")
78
V = TypeVar("V")
89

10+
NotSetType = enum.Enum("NotSetType", "VALUE")
11+
NOTSET = NotSetType.VALUE
12+
913

1014
def isinstance_noimport(obj, class_or_tuple):
1115
"""A version of isinstance that does not require importing the class.

src/parcels/_typing.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -47,6 +47,7 @@
4747
CfAxis = XgcmAxisDirection
4848
XgcmAxisPosition = Literal["center", "left", "right", "inner", "outer"]
4949
XgcmAxes = Mapping[XgcmAxisDirection, "xgcm.Axis"]
50+
VectorFields = dict[str, tuple[str, str] | tuple[str, str, str]]
5051

5152

5253
def _is_xarray_object(obj): # with no imports

0 commit comments

Comments
 (0)