Skip to content

Commit 3fba058

Browse files
committed
Move TVectorField to _typing module
1 parent a6de0ca commit 3fba058

3 files changed

Lines changed: 20 additions & 17 deletions

File tree

src/parcels/_core/fieldset.py

Lines changed: 5 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -9,19 +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
1314
from parcels._core.model import (
1415
CONSTANT_FIELD_MODELS,
1516
ModelData,
1617
StructuredModelData,
17-
TVectorField,
1818
UnstructuredModelData,
1919
)
2020
from parcels._core.utils.string import _assert_str_and_python_varname
2121
from parcels._core.utils.time import get_datetime_type_calendar
2222
from parcels._core.utils.time import is_compatible as datetime_is_compatible
2323
from parcels._python import NOTSET, NotSetType
24-
from parcels._typing import Mesh
2524
from parcels.interpolators import (
2625
XConstantField,
2726
)
@@ -151,7 +150,7 @@ def add_field(self, field: Field, name: str | None = None):
151150

152151
self.fields[name] = field
153152

154-
def add_constant_field(self, name: str, value, mesh: Mesh = "spherical"):
153+
def add_constant_field(self, name: str, value, mesh: ptyping.Mesh = "spherical"):
155154
"""Wrapper function to add a Field that is constant in space,
156155
useful e.g. when using constant horizontal diffusivity
157156
@@ -212,7 +211,7 @@ def from_ugrid_conventions(
212211
cls,
213212
ds: ux.UxDataset,
214213
mesh: str = "spherical",
215-
vector_fields: TVectorField | None | NotSetType = NOTSET,
214+
vector_fields: ptyping.VectorFields | None | NotSetType = NOTSET,
216215
):
217216
"""Create a FieldSet from a Parcels compliant uxarray.UxDataset.
218217
@@ -249,8 +248,8 @@ def from_ugrid_conventions(
249248
def from_sgrid_conventions(
250249
cls,
251250
ds: xr.Dataset,
252-
mesh: Mesh | None = None,
253-
vector_fields: TVectorField | None | NotSetType = NOTSET,
251+
mesh: ptyping.Mesh | None = None,
252+
vector_fields: ptyping.VectorFields | None | NotSetType = NOTSET,
254253
): # TODO: Update mesh to be discovered from the dataset metadata
255254
"""Create a FieldSet from a dataset using SGRID convention metadata.
256255

src/parcels/_core/model.py

Lines changed: 14 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,7 @@
99
import xarray as xr
1010

1111
import parcels._sgrid as sgrid
12+
import parcels._typing as ptyping
1213
from parcels._core.basegrid import BaseGrid
1314
from parcels._core.field import Field, VectorField
1415
from parcels._core.utils.time import TimeInterval
@@ -34,14 +35,12 @@
3435
)
3536
from parcels.interpolators._base import ScalarInterpolator, VectorInterpolator
3637

37-
TVectorField = dict[str, tuple[str, str] | tuple[str, str, str]]
38-
3938

4039
class ModelData(ABC):
4140
data: Any
4241
grid: BaseGrid
4342
field_to_interpolator: dict[str, ScalarInterpolator | VectorInterpolator]
44-
vector_field_components: TVectorField
43+
vector_field_components: ptyping.VectorFields
4544

4645
@abstractmethod
4746
def construct_fields(self) -> list[Field | VectorField]: ...
@@ -84,7 +83,7 @@ def preprocess_sgrid_model_data(ds: xr.Dataset) -> xr.Dataset:
8483

8584

8685
class StructuredModelData(ModelData):
87-
def __init__(self, data: xr.Dataset, mesh: Mesh, vector_field_components: TVectorField):
86+
def __init__(self, data: xr.Dataset, mesh: Mesh, vector_field_components: ptyping.VectorFields):
8887
if not isinstance(data, xr.Dataset):
8988
raise ValueError(f"Expected `data` to be an xarray.Dataset . Got {type(data)}")
9089

@@ -133,7 +132,7 @@ def construct_fields(self) -> list[Field | VectorField]:
133132

134133
@classmethod
135134
def from_sgrid_conventions(
136-
cls, ds: xr.Dataset, mesh: Mesh | None, vector_fields: TVectorField | None | NotSetType
135+
cls, ds: xr.Dataset, mesh: Mesh | None, vector_fields: ptyping.VectorFields | None | NotSetType
137136
) -> Self:
138137
ds = ds.copy()
139138
if mesh is None:
@@ -172,15 +171,17 @@ def from_sgrid_conventions(
172171
return model
173172

174173

175-
def resolve_vector_fields(ds: xr.Dataset, vector_fields: TVectorField | None | NotSetType) -> TVectorField:
174+
def resolve_vector_fields(
175+
ds: xr.Dataset, vector_fields: ptyping.VectorFields | None | NotSetType
176+
) -> ptyping.VectorFields:
176177
if vector_fields is None:
177178
return {}
178179
if vector_fields is NOTSET: # i.e., the default vectorfield discovery behaviour
179180
return _default_vector_field_components(list(ds.data_vars))
180181
return vector_fields
181182

182183

183-
def assert_vector_field_components_in_dataset(ds: xr.Dataset, vector_fields: TVectorField) -> None:
184+
def assert_vector_field_components_in_dataset(ds: xr.Dataset, vector_fields: ptyping.VectorFields) -> None:
184185
for components in vector_fields.values():
185186
for c in components:
186187
if c not in ds.data_vars:
@@ -220,7 +221,7 @@ def assert_vector_field_components_in_dataset(ds: xr.Dataset, vector_fields: TVe
220221

221222

222223
class UnstructuredModelData(ModelData):
223-
def __init__(self, data: ux.UxDataset, grid: UxGrid, vector_field_components: TVectorField):
224+
def __init__(self, data: ux.UxDataset, grid: UxGrid, vector_field_components: ptyping.VectorFields):
224225
if not isinstance(data, ux.UxDataset):
225226
raise ValueError(f"Expected `data` to be an uxarray.UxDataset . Got {type(data)}")
226227

@@ -259,7 +260,9 @@ def scalar_field_names(self) -> list[str]:
259260
return list(self.data.data_vars)
260261

261262
@classmethod
262-
def from_ugrid_conventions(cls, ds: ux.UxDataset, mesh: Mesh, vector_fields: TVectorField | None | NotSetType):
263+
def from_ugrid_conventions(
264+
cls, ds: ux.UxDataset, mesh: Mesh, vector_fields: ptyping.VectorFields | None | NotSetType
265+
):
263266
ds_dims = list(ds.dims)
264267
if not all(dim in ds_dims for dim in ["time", "zf", "zc"]):
265268
raise ValueError(
@@ -300,9 +303,9 @@ def _get_mesh_type_from_sgrid_dataset(ds_sgrid: xr.Dataset) -> Mesh:
300303
return "spherical" if _is_coordinate_in_degrees(ds_sgrid[fpoint_x]) else "flat"
301304

302305

303-
def _default_vector_field_components(data_vars: Sequence[Hashable]) -> TVectorField:
306+
def _default_vector_field_components(data_vars: Sequence[Hashable]) -> ptyping.VectorFields:
304307
vars = set(data_vars)
305-
ret: TVectorField = {}
308+
ret: ptyping.VectorFields = {}
306309

307310
if {"U", "V"}.issubset(vars):
308311
ret["UV"] = ("U", "V")

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)