11from __future__ import annotations
22
33from abc import ABC , abstractmethod
4+ from collections .abc import Hashable , Sequence
45from typing import Any , Self
56
67import cf_xarray # noqa: F401
78import uxarray as ux
89import xarray as xr
910
1011import parcels ._sgrid as sgrid
12+ import parcels ._typing as ptyping
1113from parcels ._core .basegrid import BaseGrid
1214from parcels ._core .field import Field , VectorField
1315from parcels ._core .utils .time import TimeInterval
1820 assert_all_field_dims_have_axis , # noqa: F401, leave import for now until decision is made # TODO v4: Make decision
1921)
2022from parcels ._logger import logger
23+ from parcels ._python import NOTSET , NotSetType
2124from parcels ._typing import Mesh
2225from parcels .convert import _ds_rename_using_standard_names
2326from 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
8185class 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+
171213CONSTANT_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
199242class 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+
279334def _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
375430def _get_time_interval (data : xr .DataArray | ux .UxDataArray ) -> TimeInterval | None :
0 commit comments