11from __future__ import annotations
22
33from abc import ABC , abstractmethod
4+ from collections .abc import Mapping , Sequence
45from typing import Any , Self
56
67import cf_xarray # noqa: F401
1819 assert_all_field_dims_have_axis , # noqa: F401, leave import for now until decision is made # TODO v4: Make decision
1920)
2021from parcels ._logger import logger
22+ from parcels ._python import _MissingType
2123from parcels ._typing import Mesh
2224from parcels .convert import _ds_rename_using_standard_names
2325from parcels .interpolators import (
3234)
3335from parcels .interpolators ._base import ScalarInterpolator , VectorInterpolator
3436
37+ TVectorFieldMapping = Mapping [str , tuple [str , str ] | tuple [str , str , str ]]
38+
3539
3640class 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
8186class 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
199208class 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+
279302def _is_coordinate_in_degrees (da : xr .DataArray ) -> bool :
280303 units = da .attrs .get ("units" )
281304 if units is None :
0 commit comments