|
9 | 9 | import xarray as xr |
10 | 10 |
|
11 | 11 | import parcels._sgrid as sgrid |
| 12 | +import parcels._typing as ptyping |
12 | 13 | from parcels._core.basegrid import BaseGrid |
13 | 14 | from parcels._core.field import Field, VectorField |
14 | 15 | from parcels._core.utils.time import TimeInterval |
|
34 | 35 | ) |
35 | 36 | from parcels.interpolators._base import ScalarInterpolator, VectorInterpolator |
36 | 37 |
|
37 | | -TVectorField = dict[str, tuple[str, str] | tuple[str, str, str]] |
38 | | - |
39 | 38 |
|
40 | 39 | class ModelData(ABC): |
41 | 40 | data: Any |
42 | 41 | grid: BaseGrid |
43 | 42 | field_to_interpolator: dict[str, ScalarInterpolator | VectorInterpolator] |
44 | | - vector_field_components: TVectorField |
| 43 | + vector_field_components: ptyping.VectorFields |
45 | 44 |
|
46 | 45 | @abstractmethod |
47 | 46 | def construct_fields(self) -> list[Field | VectorField]: ... |
@@ -84,7 +83,7 @@ def preprocess_sgrid_model_data(ds: xr.Dataset) -> xr.Dataset: |
84 | 83 |
|
85 | 84 |
|
86 | 85 | 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): |
88 | 87 | if not isinstance(data, xr.Dataset): |
89 | 88 | raise ValueError(f"Expected `data` to be an xarray.Dataset . Got {type(data)}") |
90 | 89 |
|
@@ -133,7 +132,7 @@ def construct_fields(self) -> list[Field | VectorField]: |
133 | 132 |
|
134 | 133 | @classmethod |
135 | 134 | 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 |
137 | 136 | ) -> Self: |
138 | 137 | ds = ds.copy() |
139 | 138 | if mesh is None: |
@@ -172,15 +171,17 @@ def from_sgrid_conventions( |
172 | 171 | return model |
173 | 172 |
|
174 | 173 |
|
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: |
176 | 177 | if vector_fields is None: |
177 | 178 | return {} |
178 | 179 | if vector_fields is NOTSET: # i.e., the default vectorfield discovery behaviour |
179 | 180 | return _default_vector_field_components(list(ds.data_vars)) |
180 | 181 | return vector_fields |
181 | 182 |
|
182 | 183 |
|
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: |
184 | 185 | for components in vector_fields.values(): |
185 | 186 | for c in components: |
186 | 187 | if c not in ds.data_vars: |
@@ -220,7 +221,7 @@ def assert_vector_field_components_in_dataset(ds: xr.Dataset, vector_fields: TVe |
220 | 221 |
|
221 | 222 |
|
222 | 223 | 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): |
224 | 225 | if not isinstance(data, ux.UxDataset): |
225 | 226 | raise ValueError(f"Expected `data` to be an uxarray.UxDataset . Got {type(data)}") |
226 | 227 |
|
@@ -259,7 +260,9 @@ def scalar_field_names(self) -> list[str]: |
259 | 260 | return list(self.data.data_vars) |
260 | 261 |
|
261 | 262 | @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 | + ): |
263 | 266 | ds_dims = list(ds.dims) |
264 | 267 | if not all(dim in ds_dims for dim in ["time", "zf", "zc"]): |
265 | 268 | raise ValueError( |
@@ -300,9 +303,9 @@ def _get_mesh_type_from_sgrid_dataset(ds_sgrid: xr.Dataset) -> Mesh: |
300 | 303 | return "spherical" if _is_coordinate_in_degrees(ds_sgrid[fpoint_x]) else "flat" |
301 | 304 |
|
302 | 305 |
|
303 | | -def _default_vector_field_components(data_vars: Sequence[Hashable]) -> TVectorField: |
| 306 | +def _default_vector_field_components(data_vars: Sequence[Hashable]) -> ptyping.VectorFields: |
304 | 307 | vars = set(data_vars) |
305 | | - ret: TVectorField = {} |
| 308 | + ret: ptyping.VectorFields = {} |
306 | 309 |
|
307 | 310 | if {"U", "V"}.issubset(vars): |
308 | 311 | ret["UV"] = ("U", "V") |
|
0 commit comments