Skip to content

Commit 652b85e

Browse files
committed
Update StructuredModelData.construct_fields()
1 parent 405d7f2 commit 652b85e

2 files changed

Lines changed: 37 additions & 20 deletions

File tree

src/parcels/_core/model.py

Lines changed: 36 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -19,7 +19,7 @@
1919
assert_all_field_dims_have_axis, # noqa: F401, leave import for now until decision is made # TODO v4: Make decision
2020
)
2121
from parcels._logger import logger
22-
from parcels._python import _MissingType
22+
from parcels._python import _MISSING, _MissingType
2323
from parcels._typing import Mesh
2424
from parcels.convert import _ds_rename_using_standard_names
2525
from 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+
179195
CONSTANT_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

398414
def _get_time_interval(data: xr.DataArray | ux.UxDataArray) -> TimeInterval | None:

tests/test_fieldset.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -109,6 +109,7 @@ def test_fieldset_vectorfield_default():
109109

110110
def test_fieldset_vectorfield_custom():
111111
ds1 = datasets_structured["ds_2d_left"][["U_A_grid", "V_A_grid", "grid"]].rename({"U_A_grid": "U", "V_A_grid": "V"})
112+
ds1 = ds1.rename({"U": "U_wind", "V": "V_wind"})
112113

113114
fset1 = FieldSet.from_sgrid_conventions(ds1, mesh="flat", vector_fields={"UV_wind": ("U_wind", "V_wind")})
114115

0 commit comments

Comments
 (0)