1919 assert_all_field_dims_have_axis , # noqa: F401, leave import for now until decision is made # TODO v4: Make decision
2020)
2121from parcels ._logger import logger
22- from parcels ._python import _MissingType
22+ from parcels ._python import _MISSING , _MissingType
2323from parcels ._typing import Mesh
2424from parcels .convert import _ds_rename_using_standard_names
2525from parcels .interpolators import (
@@ -116,26 +116,19 @@ def construct_fields(self) -> list[Field | VectorField]:
116116 single_fields : dict [str , Field ] = {}
117117 vector_fields : dict [str , VectorField ] = {}
118118 scalar_field_names = self .scalar_field_names
119- if "U" in scalar_field_names and "V" in scalar_field_names :
120- interp_method = XLinear_Velocity () if _is_agrid (self .data ) else CGrid_Velocity ()
121- single_fields ["U" ] = Field ("U" , self )
122- single_fields ["V" ] = Field ("V" , self )
123- vector_fields ["UV" ] = VectorField ("UV" , single_fields ["U" ], single_fields ["V" ], interp_method = interp_method )
124119
125- if "W" in scalar_field_names :
126- single_fields ["W" ] = Field ("W" , self )
127- vector_fields ["UVW" ] = VectorField (
128- "UVW" ,
129- single_fields ["U" ],
130- single_fields ["V" ],
131- single_fields ["W" ],
132- interp_method = interp_method ,
133- )
120+ for varname in set (scalar_field_names ):
121+ single_fields [varname ] = Field (str (varname ), self )
134122
135- fields : dict [str , Field | VectorField ] = {** single_fields , ** vector_fields }
136- for varname in set (scalar_field_names ) - set (fields .keys ()):
137- fields [varname ] = Field (str (varname ), self )
123+ for vfield_name , components in self .vector_field_components .items ():
124+ interp_method = (
125+ XLinear_Velocity () if _is_agrid (self .data , u = components [0 ], v = components [1 ]) else CGrid_Velocity ()
126+ )
138127
128+ component_fields = [single_fields [name ] for name in components ]
129+ vector_fields [vfield_name ] = VectorField (vfield_name , * component_fields , interp_method = interp_method )
130+
131+ fields : dict [str , Field | VectorField ] = {** single_fields , ** vector_fields }
139132 return list (fields .values ())
140133
141134 @classmethod
@@ -168,6 +161,9 @@ def from_sgrid_conventions(
168161 # ds["lon"] = ds[node_dimensions[0]]
169162 # ds["lat"] = ds[node_dimensions[1]]
170163
164+ vector_fields = resolve_vector_fields (ds , vector_fields )
165+ assert_vector_field_components_in_dataset (ds , vector_fields )
166+
171167 model = cls (ds , mesh = mesh , vector_field_components = vector_fields )
172168 model ._fields = model .construct_fields ()
173169 for f in model ._fields :
@@ -176,6 +172,26 @@ def from_sgrid_conventions(
176172 return model
177173
178174
175+ def resolve_vector_fields (
176+ ds : xr .Dataset , vector_fields : TVectorFieldMapping | None | _MissingType
177+ ) -> TVectorFieldMapping :
178+ if vector_fields is None :
179+ return {}
180+ if vector_fields is _MISSING : # i.e., the default vectorfield discovery behaviour
181+ return _default_vector_field_components (ds .data_vars )
182+ return vector_fields
183+
184+
185+ def assert_vector_field_components_in_dataset (ds : xr .Dataset , vector_fields : TVectorFieldMapping ) -> None :
186+ for components in vector_fields .values ():
187+ for c in components :
188+ if c not in ds .data_vars :
189+ raise ValueError (
190+ 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."
191+ )
192+ return
193+
194+
179195CONSTANT_FIELD_MODELS = {
180196 mesh : StructuredModelData .from_sgrid_conventions (
181197 xr .Dataset (
@@ -389,10 +405,10 @@ def _select_uxinterpolator(da: ux.UxDataArray):
389405 return None
390406
391407
392- def _is_agrid (ds : xr .Dataset ) -> bool :
408+ def _is_agrid (ds : xr .Dataset , u : str , v : str ) -> bool :
393409 # check if U and V are defined on the same dimensions
394410 # if yes, interpret as A grid
395- return set (ds ["U" ].dims ) == set (ds ["V" ].dims )
411+ return set (ds [u ].dims ) == set (ds [v ].dims )
396412
397413
398414def _get_time_interval (data : xr .DataArray | ux .UxDataArray ) -> TimeInterval | None :
0 commit comments