diff --git a/pyproject.toml b/pyproject.toml index 358e6712..18e04b2c 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -35,10 +35,9 @@ dynamic = [ "version" ] dependencies = [ "anemoi-utils>=0.5.1", "cfunits", - "earthkit-data>=0.12.4,<1", - "earthkit-geo>=0.3", - "earthkit-meteo>=0.4.1", - "earthkit-regrid>=0.4", + "earthkit-data==1.0.0rc12", + "earthkit-geo==1.0.0rc8", + "earthkit-meteo==1.0.0rc3", "healpy", "pandas<3", ] diff --git a/src/anemoi/transform/commands/get-grid.py b/src/anemoi/transform/commands/get-grid.py index 636e7b73..9d3961f3 100644 --- a/src/anemoi/transform/commands/get-grid.py +++ b/src/anemoi/transform/commands/get-grid.py @@ -47,8 +47,8 @@ def run(self, args: argparse.Namespace) -> None: else: input = args.input - ds = from_source(args.source, input) - lat, lon = ds[0].grid_points() + ds = from_source(args.source, input).to_fieldlist() + lat, lon = ds[0].geography.latlons() np.savez(args.output, latitudes=lat, longitudes=lon) diff --git a/src/anemoi/transform/commands/make-regrid-file.py b/src/anemoi/transform/commands/make-regrid-file.py index c7351def..27bc1d3b 100644 --- a/src/anemoi/transform/commands/make-regrid-file.py +++ b/src/anemoi/transform/commands/make-regrid-file.py @@ -33,8 +33,8 @@ def _ds_to_lat_lon(path: str) -> tuple[np.ndarray, np.ndarray]: import earthkit.data as ekd try: - ds = ekd.from_source("file", path) - return ds[0].grid_points() + ds = ekd.from_source("file", path).to_fieldlist() + return ds[0].geography.latlons() except TypeError: # This is a workaround for datasets that do not have data variables, # but have "latitude" and "longitude" coordinates. @@ -143,7 +143,7 @@ def run(self, args: argparse.Namespace) -> None: def make_mir_matrix(lat1, lon1, lat2, lon2, output=None, mir="mir", **mir_kwargs): import numpy as np - from earthkit.regrid.utils.mir import mir_make_matrix + from earthkit.geo.utils.mir import mir_make_matrix sparse_array = mir_make_matrix(lat1, lon1, lat2, lon2, output=None, mir=mir, **mir_kwargs) diff --git a/src/anemoi/transform/fields.py b/src/anemoi/transform/fields.py index ab7facbf..9c84c0c2 100644 --- a/src/anemoi/transform/fields.py +++ b/src/anemoi/transform/fields.py @@ -6,7 +6,7 @@ # In applying this licence, ECMWF does not waive the privileges and immunities # granted to it by virtue of its status as an intergovernmental organisation # nor does it submit to any jurisdiction. - +import datetime import logging from abc import ABC from abc import abstractmethod @@ -14,41 +14,34 @@ import earthkit.data as ekd import numpy as np -from earthkit.data.core.geography import Geography -from earthkit.data.indexing.fieldlist import SimpleFieldList - -from anemoi.transform.grids import Grid LOG = logging.getLogger(__name__) -MISSING_METADATA = object() - class Flavour(ABC): - @abstractmethod def __call__(self, key: str, field: ekd.Field) -> Any: """Called during field metadata lookup, so it can be modified""" pass -def new_fieldlist_from_list(fields: list[Any]) -> SimpleFieldList: - """Create a new SimpleFieldList from a list of fields. +def new_fieldlist_from_list(fields: list[ekd.Field]) -> ekd.FieldList: + """Create a new FieldList from a list of fields. Parameters ---------- - fields : list - List of fields to include in the FieldArray. + fields : list[ekd.Field] + List of fields to include in the fieldlist. Returns ------- - SimpleFieldList - A new SimpleFieldList containing the provided fields. + ekd.FieldList + A new FieldList containing the provided fields. """ - return SimpleFieldList(fields) + return ekd.create_fieldlist(fields) -def new_empty_fieldlist() -> SimpleFieldList: +def new_empty_fieldlist() -> ekd.FieldList: """Create a new empty SimpleFieldList. Returns @@ -56,674 +49,113 @@ def new_empty_fieldlist() -> SimpleFieldList: SimpleFieldList A new empty SimpleFieldList. """ - return SimpleFieldList([]) - - -class WrappedField: - """A wrapper around an earthkit-data field object. - - Parameters - ---------- - field : Any - The field object to wrap. - """ - - def __init__(self, field: Any) -> None: - self._field = field - - def __getattr__(self, name: str) -> Any: - """Custom attribute access method for the WrappedField class. - - Parameters - ---------- - name : str - The name of the attribute being accessed. - - Returns - ------- - Any - The value of the attribute from the underlying _field object. - - Raises - ------ - AttributeError - If the attribute name is "clone" or "copy". - """ - if name in ( - "clone", - "copy", - ): - raise AttributeError(f"{self}: forwarding of `{name}` is not supported") - - if name not in ( - "mars_area", - "mars_grid", - "to_numpy", - "metadata", - "shape", - "grid_points", - "handle", - ): - LOG.warning(f"{self}: forwarding `{name}`") - - return getattr(self._field, name) - - def __repr__(self) -> str: - """Return the string representation of the field. - - Returns - ------- - str - The string representation of the `_field` attribute. - """ - return f"{self.__class__.__name__ }({repr(self._field)}, {self._repr_specific()})" - - def _repr_specific(self) -> str: - """Return a string representation of the specific field type. - - Returns - ------- - str - The string representation of the specific field type. - """ - return f"(No specific representation for {self.__class__.__name__})" - - def clone(self, **kwargs: Any) -> "NewClonedField": - """Clone the field with new metadata. - - Parameters - ---------- - **kwargs : Any - The new metadata for the cloned field. - - Returns - ------- - NewClonedField - The cloned field with the provided metadata. - """ - return NewClonedField(self, **kwargs) - - def __iter__(self) -> Any: - """Return an iterator over the field. - - Returns - ------- - Any - An iterator over the `_field` attribute. - """ - raise NotImplementedError(f"{self}: iterating is not supported") - - -class NewDataField(WrappedField): - """Change the data of a field. - - Parameters - ---------- - field : Any - The field object to wrap. - data : np.ndarray - The new data for the field. - """ - - def __init__(self, field: Any, data: np.ndarray) -> None: - super().__init__(field) - self._data = data - self.shape = data.shape - - @property - def values(self) -> np.ndarray: - """Get the values of the field.""" - return self.to_numpy(flatten=True) - - def to_numpy(self, flatten: bool = False, dtype: type | None = None, index: Any | None = None) -> np.ndarray: - """Convert the field data to a numpy array. - - Parameters - ---------- - flatten : bool, optional - Whether to flatten the array, by default False. - dtype : type, optional - The desired data type of the array, by default None. - index : Any, optional - The index to apply to the array, by default None. - - Returns - ------- - np.ndarray - The field data as a numpy array. - """ - data = self._data - if dtype is not None: - data = data.astype(dtype) - if flatten: - data = data.flatten() - if index is not None: - data = data[index] - return data - - def _repr_specific(self) -> str: - return f"(shape={self._data.shape})" - - -class GeoMetadata(Geography): - """A wrapper around an earthkit-data Geography object. - - Parameters - ---------- - owner : Any - The owner of the geography data. - """ - - def __init__(self, owner: Any) -> None: - self.owner = owner - - def shape(self) -> tuple[int, ...]: - """Get the shape of the geography data. - - Returns - ------- - tuple - The shape of the geography data. - """ - return tuple([len(self.owner._latitudes)]) - - def resolution(self) -> str: - """Get the resolution of the geography data. - - Returns - ------- - str - The resolution of the geography data. - """ - return "unknown" - - def mars_area(self) -> list[float]: - """Get the MARS area of the geography data. - - Returns - ------- - list - The MARS area of the geography data. - """ - return [ - np.amax(self.owner._latitudes), - np.amin(self.owner._longitudes), - np.amin(self.owner._latitudes), - np.amax(self.owner._longitudes), - ] - - def mars_grid(self) -> None: - """Get the MARS grid of the geography data.""" - return None - - def latitudes(self, dtype: type | None = None) -> np.ndarray: - """Get the latitudes of the geography data. - - Parameters - ---------- - dtype : type, optional - The desired data type of the array, by default None. - - Returns - ------- - np.ndarray - The latitudes of the geography data. - """ - if dtype is None: - return self.owner._latitudes - return self.owner._latitudes.astype(dtype) - - def longitudes(self, dtype: type | None = None) -> np.ndarray: - """Get the longitudes of the geography data. - - Parameters - ---------- - dtype : type, optional - The desired data type of the array, by default None. - - Returns - ------- - np.ndarray - The longitudes of the geography data. - """ - if dtype is None: - return self.owner._longitudes - return self.owner._longitudes.astype(dtype) - - def x(self, dtype: type | None = None) -> None: - """Get the x-coordinates of the geography data.""" - raise NotImplementedError() - - def y(self, dtype: type | None = None) -> None: - """Get the y-coordinates of the geography data.""" - raise NotImplementedError() - - def _unique_grid_id(self) -> None: - """Get the unique grid ID of the geography data.""" - raise NotImplementedError() - - def projection(self) -> None: - """Get the projection of the geography data.""" - return None - - def bounding_box(self) -> None: - """Get the bounding box of the geography data.""" - raise NotImplementedError() - - def gridspec(self) -> None: - """Get the grid specification of the geography data.""" - raise NotImplementedError() - - -class NewLatLonField(WrappedField): - """Change the latitudes and longitudes of a field. - - Parameters - ---------- - field : Any - The field object to wrap. - latitudes : np.ndarray - The new latitudes for the field. - longitudes : np.ndarray - The new longitudes for the field. - """ - - def __init__(self, field: Any, latitudes: np.ndarray, longitudes: np.ndarray) -> None: - super().__init__(field) - self._latitudes = latitudes - self._longitudes = longitudes - - def grid_points(self) -> tuple[np.ndarray, np.ndarray]: - """Get the grid points of the field. - - Returns - ------- - tuple - The latitudes and longitudes of the field. - """ - return self._latitudes, self._longitudes - - def to_latlon(self, flatten: bool = True) -> dict[str, np.ndarray]: - """Convert the grid points to latitude and longitude. - - Parameters - ---------- - flatten : bool, optional - Whether to flatten the arrays, by default True. - - Returns - ------- - dict - A dictionary containing the latitudes and longitudes. - """ - assert flatten - return dict(lat=self._latitudes, lon=self._longitudes) - - def metadata(self, *args: Any, **kwargs: Any) -> Any: - """Get the metadata of the field. - - Parameters - ---------- - *args : Any - Additional arguments. - **kwargs : Any - Additional keyword arguments. - - Returns - ------- - Any - The metadata of the field. - """ - metadata = self._field.metadata(*args, **kwargs) - if hasattr(metadata, "geography"): - metadata.geography = GeoMetadata(self) - - return metadata - - -class NewGridField(WrappedField): - """Change the grid of a field. - - Parameters - ---------- - field : Any - The field object to wrap. - grid: Grid - The new grid for the field. - """ - - def __init__(self, field: Any, grid: Grid) -> None: - super().__init__(field) - self._grid = grid - - def grid_points(self) -> tuple[np.ndarray, np.ndarray]: - """Get the grid points of the field. - - Returns - ------- - tuple - The latitudes and longitudes of the field. - """ - return self._grid.latlon() - - def to_latlon(self, flatten: bool = True) -> dict[str, np.ndarray]: - """Convert the grid points to latitude and longitude. - - Parameters - ---------- - flatten : bool, optional - Whether to flatten the arrays, by default True. - - Returns - ------- - dict - A dictionary containing the latitudes and longitudes. - """ - assert flatten - coords = self._grid.latlon() - return dict(lat=coords[0], lon=coords[1]) - - def __repr__(self) -> str: - """Get the string representation of the field. - - Returns - ------- - str - The string representation of the field. - """ - return f"NewGridField({self._field}, {self._grid})" - - def metadata(self, *args: Any, **kwargs: Any) -> Any: - """Get the metadata of the field. - - Parameters - ---------- - *args : Any - Additional arguments. - **kwargs : Any - Additional keyword arguments. - - Returns - ------- - Any - The metadata of the field. - """ - metadata = self._field.metadata(*args, **kwargs) - if hasattr(metadata, "geography"): - metadata.geography = GeoMetadata(self) - - return metadata - - @property - def _latitudes(self) -> np.ndarray: - """Get the latitudes of the field.""" - return self._grid.latlon()[0] - - @property - def _longitudes(self) -> np.ndarray: - """Get the longitudes of the field.""" - return self._grid.latlon()[1] - - -class _NewMetadataField(WrappedField, ABC): - """Change the metadata of a field.""" - - def __init__(self, field: Any) -> None: - super().__init__(field) - - @abstractmethod - def mapping(self, key: str, field: ekd.Field) -> Any: ... - - def metadata(self, *args: Any, **kwargs: Any) -> Any: - """Get the metadata of the field. - - Parameters - ---------- - *args : Any - Additional arguments. - **kwargs : Any - Additional keyword arguments. - - Returns - ------- - Any - The metadata of the field. - """ - this = self - - if len(args) == 0 and len(kwargs) == 0: - - class MD: - - geography = this._field.metadata().geography - - def get(self, key, default=None): - - value = this.mapping(key, this._field) - if value is not MISSING_METADATA: - return value - - return this._field.metadata().get(key, default) - - def keys(self): - return this._field.metadata().keys() - - def __getitem__(self, key): - value = this.mapping(key, this._field) - if value is not MISSING_METADATA: - return value - - return this._field.metadata()[key] - - def override(self, *args, **kwargs): - return this._field.metadata().override(*args, **kwargs) - - return MD() - - if kwargs.get("namespace"): - assert len(args) == 0, (args, kwargs) - mars = self._field.metadata(**kwargs).copy() - for k in list(mars.keys()): - m = self.mapping(k, self._field) - if m is not MISSING_METADATA: - mars[k] = m - return mars - - def _val(a): - value = self.mapping(a, self._field) - if value is MISSING_METADATA: - return self._field.metadata(a, **kwargs) + return ekd.create_fieldlist() - if callable(value): - return value(self, a, self._field.metadata()) - return value - - result = [_val(a) for a in args] - if len(result) == 1: - return result[0] - - return tuple(result) - - -class NewMetadataField(_NewMetadataField): - """Change the metadata of a field. - - Parameters - ---------- - field : Any - The field object to wrap. - **kwargs : Any - The new metadata for the field. - """ - - def __init__(self, field: Any, **kwargs: Any) -> None: - super().__init__(field) - self.kwargs = kwargs - - def mapping(self, key: str, field: ekd.Field) -> Any: - return self.kwargs.get(key, MISSING_METADATA) - - def _repr_specific(self): - return f"(metadata={self.kwargs})" - - -class NewFlavouredField(_NewMetadataField): - def __init__(self, field: Any, flavour: Flavour) -> None: - super().__init__(field) - self.flavour = flavour - - def mapping(self, key: str, field: ekd.Field) -> Any: - return self.flavour(key, field) - - -class NewValidDateTimeField(NewMetadataField): - """Change the valid_datetime of a field. - - Parameters - ---------- - field : Any - The field object to wrap. - valid_datetime : Any - The new valid_datetime for the field. - """ - - def __init__(self, field: Any, valid_datetime: Any) -> None: - date = int(valid_datetime.strftime("%Y%m%d")) - time = int(valid_datetime.strftime("%H%M")) - - self.valid_datetime = valid_datetime - - super().__init__(field, date=date, time=time, step=0, valid_datetime=valid_datetime.isoformat()) - - -class NewClonedField(WrappedField): - """Wrapper around a field object that clones the field. - - Parameters - ---------- - field : Any - The field object to wrap. - **metadata : Any - The new metadata for the cloned field. - """ - - def __init__(self, field: Any, **metadata: Any) -> None: - super().__init__(field) - self._metadata = metadata - - def metadata(self, *args: Any, **kwargs: Any) -> Any: - """Get the metadata of the cloned field. - - Parameters - ---------- - *args : Any - Additional arguments. - **kwargs : Any - Additional keyword arguments. - - Returns - ------- - Any - The metadata of the cloned field. - """ - if len(args) == 1: - if args[0] in self._metadata: - if callable(self._metadata[args[0]]): - proc = self._metadata[args[0]] - self._metadata[args[0]] = proc(self._field, args[0], self._field.metadata()) - - if args[0] in self._metadata: - return self._metadata[args[0]] - - return self._field.metadata(*args, **kwargs) - - def _repr_specific(self): - return f"(metadata={self._metadata})" - - -def new_field_from_numpy(array: np.ndarray, *, template: WrappedField, **metadata: Any) -> NewMetadataField: +def new_field_from_numpy(array: np.ndarray, *, template: ekd.Field, **metadata: Any) -> ekd.Field: """Create a new field from a numpy array. Parameters ---------- array : np.ndarray The data for the new field. - template : WrappedField + template : ekd.Field The template field to use. **metadata : Any Additional metadata for the new field. Returns ------- - NewMetadataField + ekd.Field The new field with the provided data and metadata. """ - return NewMetadataField(NewDataField(template, array), **metadata) + new_data = template.set(**{"data.values": array}) + if not metadata: + return new_data + return new_field_with_metadata(new_data, **metadata) -def new_field_with_valid_datetime(template: WrappedField, date: Any) -> NewValidDateTimeField: - """Create a new field with a valid datetime. +def new_field_with_valid_datetime(template: ekd.Field, date: Any) -> ekd.Field: + """Create a new field with a valid datetime (sets the step to 0) + therefore updating the base_datetime as well. Parameters ---------- - template : WrappedField + template : ekd.Field The template field to use. date : Any The valid datetime for the new field. Returns ------- - NewValidDateTimeField + ekd.Field The new field with the provided valid datetime. """ - return NewValidDateTimeField(template, date) + time = template.time.set(valid_datetime=date, step=datetime.timedelta(hours=0)) + return template.set(time=time) -def new_field_with_metadata(template: WrappedField, **metadata: Any) -> NewMetadataField: +def new_field_with_metadata(template: ekd.Field, **metadata: Any) -> ekd.Field: """Create a new field with metadata. Parameters ---------- - template : WrappedField + template : ekd.Field The template field to use. **metadata : Any The metadata for the new field. Returns ------- - NewMetadataField + ekd.Field The new field with the provided metadata. """ - return NewMetadataField(template, **metadata) + key_mapping = { + "valid_datetime": "time.valid_datetime", + "base_datetime": "time.base_datetime", + "step": "time.step", + "param": "parameter.variable", + "units": "parameter.units", + "levtype": "vertical.level_type", + "levelist": "vertical.level", + "number": "ensemble.member", + } + unknown_keys = set(metadata.keys()) - set(key_mapping.keys()) + if unknown_keys: + raise ValueError(f"Unknown metadata keys: {unknown_keys}. Allowed keys are: {set(key_mapping.keys())}") -def new_field_with_units(template: WrappedField, units: str) -> NewMetadataField: + # map metadata keys to new locations + mapped_metadata = {key_mapping[key]: value for key, value in metadata.items()} + return template.set(**mapped_metadata) + + +def new_field_with_units(template: ekd.Field, units: str) -> ekd.Field: """Create a new field with units. Parameters ---------- - template : WrappedField + template : ekd.Field The template field to use. units : str The units for the new field. Returns ------- - NewMetadataField + ekd.Field The new field with the provided units. """ - return NewMetadataField(template, units=units) + return new_field_with_metadata(template, units=units) def new_field_from_latitudes_longitudes( - template: WrappedField, latitudes: np.ndarray, longitudes: np.ndarray -) -> NewLatLonField: + template: ekd.Field, latitudes: np.ndarray, longitudes: np.ndarray +) -> ekd.Field: """Create a new field from latitudes and longitudes. Parameters ---------- - template : WrappedField + template : ekd.Field The template field to use. latitudes : np.ndarray The latitudes for the new field. @@ -732,42 +164,26 @@ def new_field_from_latitudes_longitudes( Returns ------- - NewGridField + ekd.Field The new field with the provided latitudes and longitudes. """ - return NewLatLonField(template, latitudes, longitudes) - - -def new_field_from_grid( - template: WrappedField, - grid: Grid, -) -> NewGridField: - """Create a new field from a grid. - - Parameters - ---------- - template : WrappedField - The template field to use. - grid : Grid - The grid for the new field. - - Returns - ------- - NewGridField - The new field with the provided grid. - """ - return NewGridField(template, grid) + return template.set( + **{ + "geography.latitudes": latitudes, + "geography.longitudes": longitudes, + } + ) -def new_flavoured_field(field: Any, flavour: Flavour) -> NewFlavouredField: +def new_flavoured_field(field: ekd.Field, flavour: Flavour) -> ekd.Field: """Create a new field with a flavour.""" - return NewFlavouredField(field, flavour) + raise NotImplementedError("Not implemented yet.") class FieldSelection: """A class for specifying which fields to process.""" - ALLOWED_KEYS = {"param", "levelist"} + ALLOWED_KEYS = {"parameter.variable", "vertical.level"} def __init__(self, **kwargs): self._spec = kwargs @@ -792,6 +208,6 @@ def match(self, field): if self._all: return True try: - return all(field.metadata(key) in values for key, values in self._spec.items()) + return all(field.get(key) in values for key, values in self._spec.items()) except KeyError: return False diff --git a/src/anemoi/transform/filters/fields/accum_to_interval.py b/src/anemoi/transform/filters/fields/accum_to_interval.py index c6c3adbf..d0cc1b27 100644 --- a/src/anemoi/transform/filters/fields/accum_to_interval.py +++ b/src/anemoi/transform/filters/fields/accum_to_interval.py @@ -56,9 +56,9 @@ def __init__( def _identifier(self, f): # Build a unique key for time series: (name, level) - param = f.metadata("param") - level = f.metadata("level", default=None) - levelType = f.metadata("levelType", default=None) + param = f.parameter.variable() + level = f.vertical.level() + levelType = f.vertical.level_type() return (param, level, levelType) def forward(self, fields: ekd.FieldList) -> ekd.FieldList: @@ -69,7 +69,7 @@ def forward(self, fields: ekd.FieldList) -> ekd.FieldList: # Sort each group by valid time for k, fl in groups.items(): - groups[k] = sorted(fl, key=lambda x: x.metadata("valid_datetime")) + groups[k] = sorted(fl, key=lambda x: x.time.valid_datetime()) out: List[ekd.Field] = [] for (param_name, level, level_type), fl in groups.items(): diff --git a/src/anemoi/transform/filters/fields/apply_mask.py b/src/anemoi/transform/filters/fields/apply_mask.py index 2be3765c..9d7c8988 100644 --- a/src/anemoi/transform/filters/fields/apply_mask.py +++ b/src/anemoi/transform/filters/fields/apply_mask.py @@ -154,7 +154,7 @@ def prepare_filter(self): if self.path.endswith(".npy"): mask = np.load(self.path) else: - mask = ekd.from_source("file", self.path)[0].to_numpy(flatten=True) + mask = ekd.from_source("file", self.path).to_fieldlist()[0].to_numpy(flatten=True) self.mask = self._compute_mask(mask) def _compute_mask(self, mask_values: np.ndarray) -> np.ndarray: @@ -164,7 +164,7 @@ def _compute_mask(self, mask_values: np.ndarray) -> np.ndarray: def forward_select(self): if self.param is not None: - return {"param": self.param} + return {"parameter.variable": self.param} return {} def forward_transform(self, field: ekd.Field) -> ekd.Field: @@ -185,7 +185,7 @@ def forward_transform(self, field: ekd.Field) -> ekd.Field: values[self.mask] = np.nan if self.rename is not None: - param = field.metadata("param") + param = field.parameter.variable() name = f"{param}_{self.rename}" metadata["param"] = name @@ -198,7 +198,7 @@ def _separate_mask_and_fields(self, fields: ekd.FieldList) -> tuple[np.ndarray, mask_field = None remaining = [] for field in fields: - is_mask_field = field.metadata("param") == self.mask_param + is_mask_field = field.parameter.variable() == self.mask_param if is_mask_field: if mask_field is None: # store first instance of mask field diff --git a/src/anemoi/transform/filters/fields/clear_step.py b/src/anemoi/transform/filters/fields/clear_step.py index ae2f8cde..cfc056ea 100644 --- a/src/anemoi/transform/filters/fields/clear_step.py +++ b/src/anemoi/transform/filters/fields/clear_step.py @@ -8,11 +8,9 @@ # nor does it submit to any jurisdiction. -import datetime import logging import earthkit.data as ekd -from earthkit.data.utils.dates import to_datetime from anemoi.transform.fields import new_field_with_valid_datetime from anemoi.transform.fields import new_fieldlist_from_list @@ -44,8 +42,8 @@ def forward(self, data: ekd.FieldList) -> ekd.FieldList: """ result = [] for field in data: - valid_datetime = to_datetime(field.metadata("valid_datetime")) - step = field.metadata("step") - result.append(new_field_with_valid_datetime(field, valid_datetime - datetime.timedelta(hours=step))) + valid_datetime = field.time.valid_datetime() + step = field.time.step() + result.append(new_field_with_valid_datetime(field, valid_datetime - step)) return new_fieldlist_from_list(result) diff --git a/src/anemoi/transform/filters/fields/clipper.py b/src/anemoi/transform/filters/fields/clipper.py index 5074d202..26d013a0 100644 --- a/src/anemoi/transform/filters/fields/clipper.py +++ b/src/anemoi/transform/filters/fields/clipper.py @@ -62,7 +62,7 @@ def prepare_filter(self): raise ValueError("At least one value for minimum or maximum must be specified.") def forward_select(self): - return {"param": self.param} + return {"parameter.variable": self.param} def forward_transform(self, field: ekd.Field) -> ekd.Field: data = field.to_numpy() diff --git a/src/anemoi/transform/filters/fields/glacier_mask.py b/src/anemoi/transform/filters/fields/glacier_mask.py index 2ce60bed..e739200f 100644 --- a/src/anemoi/transform/filters/fields/glacier_mask.py +++ b/src/anemoi/transform/filters/fields/glacier_mask.py @@ -42,10 +42,10 @@ class SnowDepthMasked(SingleFieldFilter): optional_inputs = {"snow_depth": "sd", "snow_depth_masked": "sd_masked"} def prepare_filter(self): - self.glacier_mask = ekd.from_source("file", self.glacier_mask)[0].to_numpy().astype(bool) + self.glacier_mask = ekd.from_source("file", self.glacier_mask).to_fieldlist()[0].to_numpy().astype(bool) def forward_select(self): - return {"param": self.snow_depth} + return {"parameter.variable": self.snow_depth} def forward_transform(self, snow_depth: ekd.Field) -> ekd.Field: """Mask out glaciers in snow depth. diff --git a/src/anemoi/transform/filters/fields/icon_refinement_level.py b/src/anemoi/transform/filters/fields/icon_refinement_level.py index e5ff5acd..12aca6df 100644 --- a/src/anemoi/transform/filters/fields/icon_refinement_level.py +++ b/src/anemoi/transform/filters/fields/icon_refinement_level.py @@ -63,7 +63,7 @@ def forward(self, fields: ekd.FieldList) -> ekd.FieldList: from anemoi.utils.grids import nearest_grid_points # We assume all fields have the same grid - latitudes, longitudes = fields[0].grid_points() + latitudes, longitudes = fields[0].geography.latlons() self.nearest_grid_points = nearest_grid_points( latitudes, longitudes, diff --git a/src/anemoi/transform/filters/fields/impute_nans.py b/src/anemoi/transform/filters/fields/impute_nans.py index 5be1ec00..559f7dc9 100644 --- a/src/anemoi/transform/filters/fields/impute_nans.py +++ b/src/anemoi/transform/filters/fields/impute_nans.py @@ -47,7 +47,7 @@ class ImputeNaNs(SingleFieldFilter): required_inputs = ("param", "value") def forward_select(self): - return {"param": self.param} + return {"parameter.variable": self.param} def forward_transform(self, field: ekd.Field) -> ekd.Field: values = field.to_numpy(flatten=True).copy() diff --git a/src/anemoi/transform/filters/fields/lambda_filters.py b/src/anemoi/transform/filters/fields/lambda_filters.py index 10f96053..19b277ac 100644 --- a/src/anemoi/transform/filters/fields/lambda_filters.py +++ b/src/anemoi/transform/filters/fields/lambda_filters.py @@ -10,7 +10,7 @@ import importlib from collections.abc import Callable -from earthkit.data.core.fieldlist import Field +import earthkit.data as ekd from anemoi.transform.filter import SingleFieldFilter from anemoi.transform.filters.fields import filter_registry @@ -80,20 +80,20 @@ def prepare_filter(self): self.backward_fn = self._import_fn(self.backward_fn) def forward_select(self): - return {"param": self.param} + return {"parameter.variable": self.param} - def forward_transform(self, field: Field) -> Field: + def forward_transform(self, field: ekd.Field) -> ekd.Field: """Apply the forward lambda function to a field.""" return self.fn(field, *self.fn_args, **self.fn_kwargs) - def backward_transform(self, field: Field) -> Field: + def backward_transform(self, field: ekd.Field) -> ekd.Field: """Apply the backward lambda function to a field.""" if self.backward_fn is None: raise ValueError("Backward function is undefined.") return self.backward_fn(field, *self.fn_args, **self.fn_kwargs) @staticmethod - def _import_fn(fn: str) -> Callable[..., Field]: + def _import_fn(fn: str) -> Callable[..., ekd.Field]: """Import a function from a string path. Parameters @@ -103,7 +103,7 @@ def _import_fn(fn: str) -> Callable[..., Field]: Returns ------- - Callable[..., Field] + Callable[..., ekd.Field] The imported function. Raises diff --git a/src/anemoi/transform/filters/fields/lnsp_to_sp.py b/src/anemoi/transform/filters/fields/lnsp_to_sp.py index 502202b7..59100992 100644 --- a/src/anemoi/transform/filters/fields/lnsp_to_sp.py +++ b/src/anemoi/transform/filters/fields/lnsp_to_sp.py @@ -23,11 +23,11 @@ class LnspToSp(SingleFieldFilter): def forward_select(self): # select only fields where the param is self.log_of_surface_pressure - return {"param": self.log_of_surface_pressure} + return {"parameter.variable": self.log_of_surface_pressure} def backward_select(self): # select only fields where the param is self.surface_pressure - return {"param": self.surface_pressure} + return {"parameter.variable": self.surface_pressure} def forward_transform(self, log_of_surface_pressure: ekd.Field) -> ekd.Field: """Convert ln(sp) to sp. @@ -42,10 +42,11 @@ def forward_transform(self, log_of_surface_pressure: ekd.Field) -> ekd.Field: ekd.Field The surface pressure. """ - new_metadata = {"param": self.surface_pressure, "levelist": None, "level": None} - return self.new_field_from_numpy( + new_metadata = {"param": self.surface_pressure} + field = self.new_field_from_numpy( np.exp(log_of_surface_pressure.to_numpy()), template=log_of_surface_pressure, **new_metadata ) + return field.set(vertical={"level": None}) def backward_transform(self, surface_pressure: ekd.Field) -> ekd.Field: """Convert surface surface pressure to ln(surface_pressure). diff --git a/src/anemoi/transform/filters/fields/matching.py b/src/anemoi/transform/filters/fields/matching.py index 57255474..2c3e3633 100644 --- a/src/anemoi/transform/filters/fields/matching.py +++ b/src/anemoi/transform/filters/fields/matching.py @@ -231,12 +231,15 @@ def _transform( ekd.FieldList Transformed data. """ + if self.MATCHING.select != "param": + raise NotImplementedError("Only matching by param is supported for now.") + if self.MATCHING.vertical: grouping = GroupByParamVertical(group_by) else: grouping = GroupByParam(group_by) - - input_params = set(data.metadata(self.MATCHING.select)) + # TODO: reconsider implementation if/when fieldlist supports "parameter.variable" key + input_params = set(f.parameter.variable() for f in data) self._check_metadata_match(input_params, group_by) result: list[ekd.Field] = [] diff --git a/src/anemoi/transform/filters/fields/orog_to_z.py b/src/anemoi/transform/filters/fields/orog_to_z.py index a13be483..fc2296bc 100644 --- a/src/anemoi/transform/filters/fields/orog_to_z.py +++ b/src/anemoi/transform/filters/fields/orog_to_z.py @@ -35,11 +35,11 @@ class Orography(SingleFieldFilter): def forward_select(self): # select only fields where the param is self.orography - return {"param": self.orography} + return {"parameter.variable": self.orography} def backward_select(self): # select only fields where the param is self.geopotential - return {"param": self.geopotential} + return {"parameter.variable": self.geopotential} def forward_transform(self, orography: ekd.Field) -> ekd.Field: """Convert orography in m to surface geopotential in m²/s². diff --git a/src/anemoi/transform/filters/fields/q_height.py b/src/anemoi/transform/filters/fields/q_height.py index 0285a504..080d662f 100644 --- a/src/anemoi/transform/filters/fields/q_height.py +++ b/src/anemoi/transform/filters/fields/q_height.py @@ -45,7 +45,7 @@ def _check_consistency(A: NDArray, B: NDArray, model_level_fields: dict[str, ekd assert A.shape == B.shape, "A and B coefficients must have same shape" for name, field in model_level_fields.items(): # Assert that model levels are passed - assert all(item == "ml" for item in field.metadata("levtype")), "Field {} does not contain model levels".format( + assert all(f.vertical.level_type() == "hybrid" for f in field), "Field {} does not contain model levels".format( name, ) # Assert that A and B coefficients have one more vertical level than the model level field diff --git a/src/anemoi/transform/filters/fields/q_to_r.py b/src/anemoi/transform/filters/fields/q_to_r.py index 54f60d54..0893a58a 100644 --- a/src/anemoi/transform/filters/fields/q_to_r.py +++ b/src/anemoi/transform/filters/fields/q_to_r.py @@ -68,13 +68,13 @@ def __init__( def forward_transform(self, humidity: ekd.Field, temperature: ekd.Field) -> Iterator[ekd.Field]: """This will return the relative humidity along with temperature from specific humidity and temperature""" - pressure = 100 * float(humidity.metadata("levelist")) + pressure = 100 * float(humidity.vertical.level()) rh = thermo.relative_humidity_from_specific_humidity(temperature.to_numpy(), humidity.to_numpy(), pressure) yield self.new_field_from_numpy(rh, template=humidity, param=self.relative_humidity) def backward_transform(self, relative_humidity: ekd.Field, temperature: ekd.Field) -> Iterator[ekd.Field]: """This will return specific humidity along with temperature from relative humidity and temperature""" - pressure = 100 * float(temperature.metadata("levelist")) # levels are measured in hectopascals + pressure = 100 * float(temperature.vertical.level()) # levels are measured in hectopascals q = thermo.specific_humidity_from_relative_humidity( temperature.to_numpy(), relative_humidity.to_numpy(), pressure ) diff --git a/src/anemoi/transform/filters/fields/regrid.py b/src/anemoi/transform/filters/fields/regrid.py index d78b7263..e8f2ef11 100644 --- a/src/anemoi/transform/filters/fields/regrid.py +++ b/src/anemoi/transform/filters/fields/regrid.py @@ -14,9 +14,7 @@ import earthkit.data as ekd import numpy as np import tqdm -from earthkit.data.core.fieldlist import Field -from anemoi.transform.fields import NewLatLonField from anemoi.transform.fields import new_field_from_latitudes_longitudes from anemoi.transform.fields import new_field_from_numpy from anemoi.transform.fields import new_fieldlist_from_list @@ -48,12 +46,12 @@ def as_gridspec(grid: str | dict[str, Any] | None) -> dict[str, Any] | None: return grid -def as_griddata(grid: str | Field | dict[str, Any] | None) -> dict[str, Any] | None: +def as_griddata(grid: str | ekd.Field | dict[str, Any] | None) -> dict[str, Any] | None: """Convert grid data to a dictionary format. Parameters ---------- - grid : str | Field | dict[str, Any] | None + grid : str | ekd.Field | dict[str, Any] | None The grid data. Returns @@ -64,8 +62,8 @@ def as_griddata(grid: str | Field | dict[str, Any] | None) -> dict[str, Any] | N if grid is None: return None - if isinstance(grid, Field): - lat, lon = grid.grid_points() + if isinstance(grid, ekd.Field): + lat, lon = grid.geography.latlons() return dict(latitudes=lat, longitudes=lon) if isinstance(grid, dict) and "latitudes" in grid and "longitudes" in grid: @@ -91,11 +89,11 @@ class RegridFilter(Filter): When building a dataset for a specific model, it is possible that the source grid or resolution does not fit the needs. In that case, it is possible to add a filter to interpolate the data to a target grid. It - will call the ``interpolate`` function from `earthkit-regrid - `_ if + will call the ``regrid`` function from `earthkit-geo + `_ if the keys ``method``, ``in_grid`` and ``out_grid`` are provided and if a `pre-generated matrix - `_ + `_ exists for this transformation. Otherwise, it is possible to provide a ``regrid matrix`` previously generated with :ref:`make-regrid-file`. The generated matrix is an NPZ file containing the @@ -230,7 +228,7 @@ def __init__(self, *, in_grid: Any, out_grid: Any, method: str = "linear", check if check: LOG.warning("Check is not supported by EarthkitRegrid") - def __call__(self, field: Any) -> NewLatLonField: + def __call__(self, field: Any) -> ekd.Field: """Interpolate the field data. Parameters @@ -240,19 +238,23 @@ def __call__(self, field: Any) -> NewLatLonField: Returns ------- - NewLatLonField + ekd.Field The interpolated field. """ - from earthkit.regrid import interpolate + from earthkit.geo.grids.array import regrid + + regrid_result = regrid( + field.to_numpy(flatten=True), + in_grid=self.in_grid, + out_grid=self.out_grid, + interpolation=self.method, + ) + # regrid returns (data, grid_spec) + regrid_data, _ = regrid_result return new_field_from_latitudes_longitudes( new_field_from_numpy( - interpolate( - field.to_numpy(flatten=True), - in_grid=self.in_grid, - out_grid=self.out_grid, - method=self.method, - ), + regrid_data, template=field, ), **self.out_griddata, @@ -289,17 +291,17 @@ def __init__(self, *, matrix: str, check: bool) -> None: latitudes=loaded["out_latitudes"], longitudes=loaded["out_longitudes"] ) - def __call__(self, field: Field) -> NewLatLonField: + def __call__(self, field: ekd.Field) -> ekd.Field: """Interpolate the field data using the regrid matrix. Parameters ---------- - field : Field + field : ekd.Field The field to be interpolated. Returns ------- - NewLatLonField + ekd.Field The interpolated field. """ if self.check: @@ -343,7 +345,7 @@ def __init__(self, *, in_grid: Any, out_grid: Any, method: str, check: bool = Fa if check: LOG.warning("Check is not supported by ScipyKDTreeNearestNeighbours") - def __call__(self, field: Any) -> NewLatLonField: + def __call__(self, field: Any) -> ekd.Field: """Interpolate the field data using nearest neighbours. Parameters @@ -353,7 +355,7 @@ def __call__(self, field: Any) -> NewLatLonField: Returns ------- - NewLatLonField + ekd.Field The interpolated field. """ if self.in_grid is None: @@ -401,7 +403,7 @@ def __init__(self, *, mask: str, check: bool) -> None: self.mask = np.load(mask)["mask"] - def __call__(self, field: Field) -> NewLatLonField: + def __call__(self, field: ekd.Field) -> ekd.Field: """Regrid the field data using the mask. Parameters @@ -411,7 +413,7 @@ def __call__(self, field: Field) -> NewLatLonField: Returns ------- - NewLatLonField + ekd.Field The regridded field. """ @@ -420,7 +422,7 @@ def __call__(self, field: Field) -> NewLatLonField: data = data[..., self.mask] if self.out_latitudes is None or self.out_longitudes is None: - in_latitudes, in_longitudes = field.grid_points() + in_latitudes, in_longitudes = field.geography.latlons() self.out_latitudes = in_latitudes[self.mask] self.out_longitudes = in_longitudes[self.mask] diff --git a/src/anemoi/transform/filters/fields/remove_nans.py b/src/anemoi/transform/filters/fields/remove_nans.py index 907e3e5a..9ab41a7e 100644 --- a/src/anemoi/transform/filters/fields/remove_nans.py +++ b/src/anemoi/transform/filters/fields/remove_nans.py @@ -92,7 +92,7 @@ def forward(self, fields: ekd.FieldList) -> ekd.FieldList: first = fields[0] else: for first in fields: - if first.metadata("param") == self.param: + if first.parameter.variable() == self.param: break else: raise ValueError(f"{self.param=} not found in\n{fields.ls}") @@ -100,9 +100,9 @@ def forward(self, fields: ekd.FieldList) -> ekd.FieldList: data = first.to_numpy(flatten=True) self._mask = ~np.isnan(data) - latitudes, longitudes = first.grid_points() - self._latitudes = latitudes[self._mask] - self._longitudes = longitudes[self._mask] + latitudes, longitudes = first.geography.latlons() + self._latitudes = latitudes.flatten()[self._mask] + self._longitudes = longitudes.flatten()[self._mask] result = [] for field in tqdm.tqdm(fields, desc="Remove NaNs"): diff --git a/src/anemoi/transform/filters/fields/rename.py b/src/anemoi/transform/filters/fields/rename.py index 23658baf..7ac724c3 100644 --- a/src/anemoi/transform/filters/fields/rename.py +++ b/src/anemoi/transform/filters/fields/rename.py @@ -15,6 +15,32 @@ from anemoi.transform.filter import SingleFieldFilter from anemoi.transform.filters.fields import filter_registry +# Mapping from old metadata keys to component-based accessor paths +_KEY_MAPPING = { + "param": "parameter.variable", + "levelist": "vertical.level", + "levtype": "vertical.level_type", + "step": "time.step", + "valid_datetime": "time.valid_datetime", + "number": "ensemble.member", +} + + +def _get_metadata(field, key): + """Get metadata value by key, trying original metadata keys first, then component API.""" + try: + return field.metadata(key) + except (KeyError, TypeError): + pass + + # Try the mapped component key + mapped = _KEY_MAPPING.get(key) + if mapped is not None: + try: + return field.get(mapped) + except (KeyError, TypeError) as e: + raise KeyError(f"Cannot get metadata for key '{key}'") from e + class FormatRename: def __init__(self, what, format): @@ -28,18 +54,14 @@ def __init__(self, what, format): self.format_keys = [b.replace(":", self._delimiter) for b in self.bits] def rename(self, field): - md = field.metadata(self.what, default=None) + try: + md = _get_metadata(field, self.what) + except KeyError: + return field if md is None: return field - values = field.metadata(*self.bits) - values = ( - [ - values, - ] - if isinstance(values, str) - else values - ) + values = [_get_metadata(field, b) for b in self.bits] kwargs = dict(zip(self.format_keys, values)) kwargs = {self.what: self.format.format(**kwargs)} @@ -52,7 +74,10 @@ def __init__(self, what, renaming): self.renaming = renaming def rename(self, field): - md = field.metadata(self.what, default=None) + try: + md = _get_metadata(field, self.what) + except KeyError: + return field if md is None: return field diff --git a/src/anemoi/transform/filters/fields/rescale.py b/src/anemoi/transform/filters/fields/rescale.py index 65f81aae..e5d7eb6e 100644 --- a/src/anemoi/transform/filters/fields/rescale.py +++ b/src/anemoi/transform/filters/fields/rescale.py @@ -44,7 +44,7 @@ def prepare_filter(self): raise NotImplementedError("prepare_filter must be implemented by subclasses.") def forward_select(self): - return {"param": self.param} + return {"parameter.variable": self.param} def forward_transform(self, param: ekd.Field) -> ekd.Field: """Apply the forward transformation (x to ax+b).""" diff --git a/src/anemoi/transform/filters/fields/rotate_winds.py b/src/anemoi/transform/filters/fields/rotate_winds.py index 019829d1..3027940e 100644 --- a/src/anemoi/transform/filters/fields/rotate_winds.py +++ b/src/anemoi/transform/filters/fields/rotate_winds.py @@ -71,20 +71,24 @@ def forward_transform(self, x_wind: ekd.Field, y_wind: ekd.Field) -> Iterator[ek Iterator[ekd.Field] The rotated wind component fields. """ - lats, lons = x_wind.grid_points() - proj_string = str(x_wind.projection()) + lats, lons = x_wind.geography.latlons() + if self.source_projection is not None: + source_proj = self.source_projection + else: + projection = x_wind.geography.projection() + source_proj = CRS.from_string(projection.to_proj_string()) x_new, y_new = rotate_vector( lats, lons, - x_wind.to_numpy(flatten=True), - y_wind.to_numpy(flatten=True), - (self.source_projection if self.source_projection is not None else CRS.from_string(proj_string)), + x_wind.to_numpy(), + y_wind.to_numpy(), + source_proj, self.target_projection, ) - yield self.new_field_from_numpy(x_new, template=x_wind, param=x_wind.metadata("param")) - yield self.new_field_from_numpy(y_new, template=y_wind, param=y_wind.metadata("param")) + yield self.new_field_from_numpy(x_new, template=x_wind, param=x_wind.parameter.variable()) + yield self.new_field_from_numpy(y_new, template=y_wind, param=y_wind.parameter.variable()) def backward_transform(self, x_wind: ekd.Field, y_wind: ekd.Field) -> Iterator[ekd.Field]: """Rotate wind components from target projection back to source projection. @@ -101,21 +105,21 @@ def backward_transform(self, x_wind: ekd.Field, y_wind: ekd.Field) -> Iterator[e Iterator[ekd.Field] The rotated wind component fields. """ - lats, lons = x_wind.grid_points() + lats, lons = x_wind.geography.latlons() assert self.source_projection is not None, "source_projection cannot be None when unrotating winds!" x_unrotated, y_unrotated = rotate_vector( lats, lons, - x_wind.to_numpy(flatten=True), - y_wind.to_numpy(flatten=True), + x_wind.to_numpy(), + y_wind.to_numpy(), self.target_projection, self.source_projection, ) - yield self.new_field_from_numpy(x_unrotated, template=x_wind, param=x_wind.metadata("param")) - yield self.new_field_from_numpy(y_unrotated, template=y_wind, param=y_wind.metadata("param")) + yield self.new_field_from_numpy(x_unrotated, template=x_wind, param=x_wind.parameter.variable()) + yield self.new_field_from_numpy(y_unrotated, template=y_wind, param=y_wind.parameter.variable()) filter_registry.register("rotate_winds", RotateWinds) diff --git a/src/anemoi/transform/filters/fields/sum.py b/src/anemoi/transform/filters/fields/sum.py index 71e51df9..1a13fc15 100644 --- a/src/anemoi/transform/filters/fields/sum.py +++ b/src/anemoi/transform/filters/fields/sum.py @@ -18,6 +18,7 @@ from anemoi.transform.fields import new_fieldlist_from_list from anemoi.transform.filter import Filter from anemoi.transform.filters.fields import filter_registry +from anemoi.transform.grouping import grouping_dict_all LOG = logging.getLogger(__name__) @@ -84,16 +85,16 @@ def forward(self, fields: ekd.FieldList) -> ekd.FieldList: needed_fields: dict[tuple[Hashable, ...], dict[str, ekd.Field]] = defaultdict(dict) for f in fields: - key = f.metadata(namespace="mars") - param = key.pop("param", None) + key = grouping_dict_all(f) + param = key.pop("parameter.variable") if self.ignore_level: - ll = key.pop("levelist", None) + ll = key.pop("vertical.level", None) LOG.debug(f"Removing levelist ({ll}) from matching key for variable: {param}") if param is None: - param = f.metadata("param") + param = f.parameter.variable() if param in self.params: - key = tuple(key.items()) + key = frozenset(key.items()) if param in needed_fields[key]: raise ValueError(f"Duplicate field {param} for {key}") diff --git a/src/anemoi/transform/filters/fields/timeseries.py b/src/anemoi/transform/filters/fields/timeseries.py index 462d3a3f..a0c58c93 100644 --- a/src/anemoi/transform/filters/fields/timeseries.py +++ b/src/anemoi/transform/filters/fields/timeseries.py @@ -68,7 +68,7 @@ def forward_transform(self, template_param: ekd.Field) -> Iterator[ekd.Field]: Iterator[ekd.Field] Transformed fields. """ - dt = template_param.metadata("valid_datetime") + dt = template_param.time.valid_datetime() template_array = template_param.to_numpy() sel = self.ds.sel(time=dt) diff --git a/src/anemoi/transform/flavour.py b/src/anemoi/transform/flavour.py index 5562cd88..391c39f5 100644 --- a/src/anemoi/transform/flavour.py +++ b/src/anemoi/transform/flavour.py @@ -20,6 +20,30 @@ from anemoi.transform.fields import new_flavoured_field +class _FieldMetadataMapping: + """A Mapping-like wrapper around a field that supports key lookup via field.metadata(key). + + This is used to provide a dict-like interface for Rule.match(), which expects + a Mapping with __contains__ and __getitem__. + """ + + def __init__(self, field: ekd.Field) -> None: + self._field = field + + def __contains__(self, key: str) -> bool: + try: + self._field.metadata(key) + return True + except (KeyError, TypeError): + return False + + def __getitem__(self, key: str) -> Any: + try: + return self._field.metadata(key) + except (KeyError, TypeError): + raise KeyError(key) + + class RuleBasedFlavour(Flavour): """Rule-based flavour for GRIB files.""" @@ -93,7 +117,7 @@ def __call__(self, key: str, field: ekd.Field) -> Any: return MISSING_METADATA for rule in self.rules[key]: - if rule.match(field.metadata()): + if rule.match(_FieldMetadataMapping(field)): return rule.result return MISSING_METADATA diff --git a/src/anemoi/transform/grids/unstructured.py b/src/anemoi/transform/grids/unstructured.py index 0095b3fd..11f3262d 100644 --- a/src/anemoi/transform/grids/unstructured.py +++ b/src/anemoi/transform/grids/unstructured.py @@ -13,8 +13,8 @@ from urllib.parse import urlparse import numpy as np +from earthkit.data import SimpleFieldList from earthkit.data import from_source -from earthkit.data.indexing.fieldlist import FieldArray LOG = logging.getLogger(__name__) @@ -74,7 +74,7 @@ def _load(url_or_path: str, param: str) -> tuple[np.ndarray, str]: else: source = "file" - ds = from_source(source, url_or_path) + ds = from_source(source, url_or_path).to_fieldlist() ds = ds.sel(param=param) assert len(ds) == 1, f"{url_or_path} {param}, expected one field, got {len(ds)}" @@ -153,7 +153,7 @@ def to_latlon(self, flatten: bool = False) -> dict[str, np.ndarray]: return dict(lat=self.geography.latitudes, lon=self.geography.longitudes) -class UnstructuredGridFieldList(FieldArray): +class UnstructuredGridFieldList(SimpleFieldList): """List of unstructured grid fields.""" @classmethod diff --git a/src/anemoi/transform/grouping/__init__.py b/src/anemoi/transform/grouping/__init__.py index d15ee5d3..63740162 100644 --- a/src/anemoi/transform/grouping/__init__.py +++ b/src/anemoi/transform/grouping/__init__.py @@ -14,7 +14,7 @@ from collections.abc import Iterator from typing import Any -from earthkit.data import SimpleFieldList +import earthkit.data as ekd LOG = logging.getLogger(__name__) @@ -52,6 +52,23 @@ def _flatten(params: list[Any] | tuple[Any, ...]) -> list[str]: return flat +def grouping_dict_all(field): + # replacement for field.metadata(namespace='mars') non-specific grouping + KEYS = ( + "parameter.variable", + "time.valid_datetime", + "time.step", + "ensemble.member", + "vertical.level", + "vertical.level_type", + ) + result = dict(zip(KEYS, field.get(KEYS, default=None))) + # Normalise missing ensemble member to "0" to prevent grouping inconsistencies + if result["ensemble.member"] is None: + result["ensemble.member"] = "0" + return result + + class GroupByParam: """Group matching fields by parameters name. @@ -67,24 +84,12 @@ def __init__(self, params: list[str]) -> None: self.params = _flatten(params) @staticmethod - def _get_grouping_key( - field, extract_from_grouping_key: list[str], remove_from_grouping_key: list[str] | None = None - ): - remove_from_grouping_key = remove_from_grouping_key or [] - - grouping_key = field.metadata(namespace="mars") - if not grouping_key: - meta_keys = [k for k in field.metadata().keys() if k not in ("latitudes", "longitudes", "values")] - grouping_key = {k: field.metadata(k) for k in meta_keys} - if not meta_keys: - raise NotImplementedError(f"GroupByParam: {field} has no sufficient metadata") + def _get_grouping_key(field, extract_from_grouping_key: list[str]): + grouping_key = grouping_dict_all(field) extracted_keys = {} for key in extract_from_grouping_key: - extracted_keys[key] = grouping_key.pop(key, field.metadata().get(key, default=None)) - - for key in remove_from_grouping_key: - grouping_key.pop(key, None) + extracted_keys[key] = grouping_key.pop(key) if len(extracted_keys) != len(extract_from_grouping_key): raise ValueError(f"Expected {extract_from_grouping_key} keys to extract, got {extracted_keys}") @@ -95,10 +100,8 @@ def _get_groups(self, data: list[Any], *, other: Callable[[Any], None] = _lost) self.groups: dict[frozenset[Any], dict[str, Any]] = defaultdict(dict) self.groups_params = set() for f in data: - key, extras = self._get_grouping_key( - f, extract_from_grouping_key=["param"], remove_from_grouping_key=["variable"] - ) - param = extras["param"] + key, extras = self._get_grouping_key(f, extract_from_grouping_key=["parameter.variable"]) + param = extras["parameter.variable"] if param not in self.params: other(f) @@ -145,10 +148,11 @@ def _get_groups(self, data: list[Any], *, other: Callable[[Any], None] = _lost) levels: dict[str, Any] = defaultdict(list) for f in data: key, extras = self._get_grouping_key( - f, extract_from_grouping_key=["param", "levelist"], remove_from_grouping_key=["variable", "levtype"] + f, extract_from_grouping_key=["parameter.variable", "vertical.level", "vertical.level_type"] ) - param = extras["param"] - level = extras["levelist"] + param = extras["parameter.variable"] + level = extras["vertical.level"] + level_type = extras["vertical.level_type"] if param not in self.params: other(f) @@ -156,7 +160,7 @@ def _get_groups(self, data: list[Any], *, other: Callable[[Any], None] = _lost) key = frozenset(key.items()) - if level is None: + if level is None or level_type != "hybrid": if param in self.groups[key]: raise ValueError(f"Duplicate component {param} for {key}") self.groups[key][param] = f @@ -167,9 +171,14 @@ def _get_groups(self, data: list[Any], *, other: Callable[[Any], None] = _lost) else: self.groups[key][param].append(f) else: - ds = SimpleFieldList() - ds.append(f) - self.groups[key][param] = ds + self.groups[key][param] = [f] levels[param].append(level) self.groups_params.add(param) + + # Convert accumulated lists to FieldLists + for key, group in self.groups.items(): + for param, value in group.items(): + if isinstance(value, list): + group[param] = ekd.create_fieldlist(value) + LOG.info(f"Params groups: {self.groups_params}") diff --git a/src/anemoi/transform/sources/mars.py b/src/anemoi/transform/sources/mars.py index e98df4f1..855a1b76 100644 --- a/src/anemoi/transform/sources/mars.py +++ b/src/anemoi/transform/sources/mars.py @@ -43,7 +43,7 @@ def forward(self, data: dict[str, Any]) -> ekd.Source: ekd.Source The data fetched from MARS. """ - return ekd.from_source("mars", **data) + return ekd.from_source("mars", **data).to_fieldlist() def __ror__(self, data: dict[str, Any]) -> Source: """Enable the use of the pipe operator with this source. diff --git a/src/anemoi/transform/variables/__init__.py b/src/anemoi/transform/variables/__init__.py index a171cd94..da2569f4 100644 --- a/src/anemoi/transform/variables/__init__.py +++ b/src/anemoi/transform/variables/__init__.py @@ -69,7 +69,7 @@ def from_earthkit(cls, name: str, field: Any) -> Any: Any The created Variable instance. """ - from anemoi.transform.variables.from_ekd import VariableFromEarthkit + from anemoi.transform.variables.from_dict import VariableFromEarthkit return VariableFromEarthkit(name, field) diff --git a/src/anemoi/transform/variables/from_dict.py b/src/anemoi/transform/variables/from_dict.py index 13cb626b..6a4841f4 100644 --- a/src/anemoi/transform/variables/from_dict.py +++ b/src/anemoi/transform/variables/from_dict.py @@ -170,6 +170,101 @@ def __init__(self, name: str, data: dict[str, Any]) -> None: super().__init__(name, data) +class VariableFromEarthkit(VariableFromMarsVocabulary): + """A variable that is defined by an EarthKit field.""" + + # Mapping from original metadata keys to earthkit component accessors + _MARS_KEY_MAPPING = { + "param": "parameter.variable", + "levtype": "vertical.level_type", + "levelist": "vertical.level", + "step": "time.step", + "number": "ensemble.member", + } + + # Mapping from earthkit 1.0 level type names to MARS-style abbreviations + _LEVEL_TYPE_MAPPING = { + "surface": "sfc", + "pressure": "pl", + "model": "ml", + "depth_below_ground_level": "sfc", + "height_above_ground": "sfc", + "potential_vorticity": "pv", + "potential_temperature": "pt", + } + + def __init__(self, name: str, field: Any) -> None: + """Initialize the variable with a name and field. + + Parameters + ---------- + name : str + The name of the variable. + field : Any + The EarthKit field defining the variable. + """ + # Build a MARS-like metadata dict from the field's component API + mars_data = {} + for mars_key, component_key in self._MARS_KEY_MAPPING.items(): + try: + mars_data[mars_key] = field.get(component_key) + except (KeyError, TypeError): + pass + mars_data["param"] = name + + data = {"mars": mars_data} + + # Convert earthkit level type to MARS-style abbreviation + if "levtype" in mars_data: + levtype = mars_data["levtype"] + if levtype in self._LEVEL_TYPE_MAPPING: + mars_data["levtype"] = self._LEVEL_TYPE_MAPPING[levtype] + else: + # Remove unknown/unmapped level types so they are treated as None + del mars_data["levtype"] + + # Get units from the field if available + try: + units = field.get("parameter.units") + if units is not None: + data["units"] = str(units) + except (KeyError, TypeError): + pass + + # Try to extract time processing info from the field + try: + statistical_process = field.get("time.statistical_process") + if statistical_process is not None: + data["process"] = statistical_process + except (KeyError, TypeError): + pass + + super().__init__(name, data) + self.field = field + # Track whether we actually got process info from the field + self._has_process_info = "process" in data + + @property + def is_instantanous(self) -> bool: + """Check if the variable is instantaneous. + + Returns None if this information is not available from the field. + """ + if not self._has_process_info: + return None + return super().is_instantanous + + @property + def period(self): + """Get the variable's period. + + Returns None if time processing info is not available from the field. + """ + if not self._has_process_info: + return None + return super().period + + class PostProcessedVariable(VariableFromMarsVocabulary): """A variable that is defined by a post-processed dictionary.""" diff --git a/tests/conftest.py b/tests/conftest.py index f701c9a0..4d5a4eb3 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -14,27 +14,13 @@ import numpy as np import pytest from anemoi.utils.testing import GetTestData -from earthkit.data.indexing.fieldlist import SimpleFieldList -from earthkit.data.sources.array_list import ArrayField -from earthkit.data.utils.metadata.dict import UserMetadata from anemoi.transform.source import Source from anemoi.transform.sources import source_registry -pytest_plugins = ["anemoi.utils.testing"] - -# Create a ekd Metadata Class that mocks the mars metadata namespace -MARS_KEYS = {"param", "levelist", "type", "step", "date", "time", "number", "expver", "class", "stream", "domain"} - +from .utils import group_component_dict -class MarsUserMetadata(UserMetadata): - def namespaces(self): - return ["mars"] - - def as_namespace(self, namespace=None): - if namespace == "mars": - return {k: v for k, v in self._data.items() if k in MARS_KEYS} - return {} +pytest_plugins = ["anemoi.utils.testing"] @source_registry.register("testing") @@ -50,7 +36,7 @@ def forward(self, *args, **kwargs): @pytest.fixture def fieldlist(get_test_data: GetTestData) -> ekd.FieldList: """Fixture to create a fieldlist for testing.""" - return ekd.from_source("file", get_test_data("anemoi-filters/2t-sp.grib")) + return ekd.from_source("file", get_test_data("anemoi-filters/2t-sp.grib")).to_fieldlist() @pytest.fixture @@ -58,9 +44,27 @@ def test_source(get_test_data: GetTestData) -> Callable[[str | list[dict]], Sour def _source(dataset: str | list[dict]) -> Source: """Create a source from a known file or a list of dicts for testing.""" if isinstance(dataset, str): - ds = ekd.from_source("file", get_test_data(dataset)) + path = get_test_data(dataset) + # TODO: revisit how npy files are loaded + if path.endswith(".npy"): + # numpy files can't be loaded by earthkit 1.0, load directly + arr = np.load(path) + + class _NumpyWrapper: + """Wrapper to mimic the old ds interface for numpy arrays.""" + + def to_numpy(self): + return arr + + def to_fieldlist(self): + return self + + ds = _NumpyWrapper() + else: + ds = ekd.from_source("file", path).to_fieldlist() elif isinstance(dataset, list): - ds = ekd.from_source("list-of-dicts", dataset) + dataset = [group_component_dict(spec) for spec in dataset] + ds = ekd.from_source("list-of-dicts", dataset).to_fieldlist() else: raise ValueError("dataset must be a string or a list of dicts") return source_registry.create("testing", dataset=ds) @@ -71,10 +75,8 @@ def _source(dataset: str | list[dict]) -> Source: @pytest.fixture def mars_test_source() -> Callable[[list[dict]], Source]: def _source(dataset: list[dict]) -> Source: - fields = [] - for d in dataset: - v = np.array(d["values"]) - fields.append(ArrayField(v, MarsUserMetadata(d, shape=v.shape))) - return source_registry.create("testing", dataset=SimpleFieldList(fields=fields)) + dataset = [group_component_dict(spec) for spec in dataset] + ds = ekd.from_source("list-of-dicts", dataset).to_fieldlist() + return source_registry.create("testing", dataset=ds) return _source diff --git a/tests/dispatching_filters/test_clip.py b/tests/dispatching_filters/test_clip.py index 40a39998..a2669081 100644 --- a/tests/dispatching_filters/test_clip.py +++ b/tests/dispatching_filters/test_clip.py @@ -20,7 +20,7 @@ def calc_stats(fieldlist): stats = {} for param in ("2t", "sp"): - fields = fieldlist.sel(param=param) + fields = fieldlist.sel(**{"parameter.variable": param}) assert len(fields) == 1 data = fields[0].to_numpy() stats[param] = {"min": np.min(data), "max": np.max(data)} diff --git a/tests/dispatching_filters/test_geopotential_to_height.py b/tests/dispatching_filters/test_geopotential_to_height.py index f296e970..84e9a3d0 100644 --- a/tests/dispatching_filters/test_geopotential_to_height.py +++ b/tests/dispatching_filters/test_geopotential_to_height.py @@ -17,9 +17,9 @@ from ..utils import collect_fields_by_param MOCK_FIELD_METADATA = { - "latitudes": [10.0, 0.0, -10.0], - "longitudes": [20, 40.0], - "valid_datetime": "2018-08-01T09:00:00Z", + "geography.distinct_latitudes": [10.0, 0.0, -10.0], + "geography.distinct_longitudes": [20, 40.0], + "time.valid_datetime": "2018-08-01T09:00:00Z", } OROG_VALUES = np.array([[243.87788459, 1892.45371246], [427.80215359, 156.92873391], [2167.93458212, 338.15794671]]) @@ -27,7 +27,7 @@ @pytest.fixture def orog_source(test_source): - OROG_SPEC = [{"param": "orog", "values": OROG_VALUES, **MOCK_FIELD_METADATA}] + OROG_SPEC = [{"parameter.variable": "orog", "data.values": OROG_VALUES, **MOCK_FIELD_METADATA}] return test_source(OROG_SPEC) diff --git a/tests/dispatching_filters/test_impute_nans.py b/tests/dispatching_filters/test_impute_nans.py index 01286c74..67169e5f 100644 --- a/tests/dispatching_filters/test_impute_nans.py +++ b/tests/dispatching_filters/test_impute_nans.py @@ -16,9 +16,9 @@ from ..utils import collect_fields_by_param INPUT_METADATA = { - "latitudes": [10.0, 0.0, -10.0], - "longitudes": [20.0, 30.0, 40.0], - "valid_datetime": "2018-08-01T12:00:00Z", + "geography.distinct_latitudes": [10.0, 0.0, -10.0], + "geography.distinct_longitudes": [20.0, 30.0, 40.0], + "time.valid_datetime": "2018-08-01T12:00:00Z", } T_VALUES = np.array([[1.0, np.nan, 3.0], [np.nan, 5.0, 6.0], [7.0, np.nan, 9.0]]) @@ -29,8 +29,8 @@ def source(test_source): return test_source( [ - {"param": "t", "values": T_VALUES.copy(), **INPUT_METADATA}, - {"param": "q", "values": Q_VALUES.copy(), **INPUT_METADATA}, + {"parameter.variable": "t", "data.values": T_VALUES.copy(), **INPUT_METADATA}, + {"parameter.variable": "q", "data.values": Q_VALUES.copy(), **INPUT_METADATA}, ] ) diff --git a/tests/dispatching_filters/test_mask.py b/tests/dispatching_filters/test_mask.py index d4310c60..5b7783e4 100644 --- a/tests/dispatching_filters/test_mask.py +++ b/tests/dispatching_filters/test_mask.py @@ -18,9 +18,9 @@ from ..utils import collect_fields_by_param MOCK_FIELD_METADATA = { - "latitudes": [10.0, 0.0, -10.0], - "longitudes": [20, 40.0], - "valid_datetime": "2018-08-01T09:00:00Z", + "geography.distinct_latitudes": [10.0, 0.0, -10.0], + "geography.distinct_longitudes": [20, 40.0], + "time.valid_datetime": "2018-08-01T09:00:00Z", } MASK_VALUES = { @@ -40,7 +40,8 @@ @pytest.fixture() def field_source(test_source): FIELD_SPECS = [ - {"param": param, "values": values.copy(), **MOCK_FIELD_METADATA} for param, values in DATA_VALUES.items() + {"parameter.variable": param, "data.values": values.copy(), **MOCK_FIELD_METADATA} + for param, values in DATA_VALUES.items() ] return test_source(FIELD_SPECS) @@ -54,7 +55,12 @@ def side_effect(source_type, path): # mask expected to be flattened mask = MASK_VALUES[path].copy().flatten() mock_field.to_numpy.return_value = mask - return [mock_field] + # Return a mock that supports .to_fieldlist()[0] + mock_source = mock.Mock() + mock_fieldlist = mock.Mock() + mock_fieldlist.__getitem__ = mock.Mock(return_value=mock_field) + mock_source.to_fieldlist.return_value = mock_fieldlist + return mock_source with mock.patch("anemoi.transform.filters.fields.apply_mask.ekd.from_source", autospec=True) as mock_fn: mock_fn.side_effect = side_effect diff --git a/tests/dispatching_filters/test_remove_nans.py b/tests/dispatching_filters/test_remove_nans.py index 6f82b85c..f88fc9c1 100644 --- a/tests/dispatching_filters/test_remove_nans.py +++ b/tests/dispatching_filters/test_remove_nans.py @@ -16,9 +16,9 @@ from ..utils import collect_fields_by_param INPUT_METADATA = { - "latitudes": [10.0, 0.0, -10.0], - "longitudes": [20.0, 30.0, 40.0], - "valid_datetime": "2018-08-01T12:00:00Z", + "geography.distinct_latitudes": [10.0, 0.0, -10.0], + "geography.distinct_longitudes": [20.0, 30.0, 40.0], + "time.valid_datetime": "2018-08-01T12:00:00Z", } INPUT_VALUES = [ @@ -38,17 +38,18 @@ EXPECTED_METADATA = { # take the original (flattened) versions and remove where there were NaNs in the first field - # "latitudes": [10.0, ---, 10.0, ---, 0.0, ---, -10.0, -10.0, ---], - "latitudes": [10.0, 10.0, 0.0, -10.0, -10.0], - # "longitudes": [20.0, ---, 40.0, ---, 30.0, ---, 20.0, 30.0, ---], - "longitudes": [20.0, 40.0, 30.0, 20.0, 30.0], + # "geography.latitudes": [10.0, ---, 10.0, ---, 0.0, ---, -10.0, -10.0, ---], + "geography.latitudes": [10.0, 10.0, 0.0, -10.0, -10.0], + # "geography.longitudes": [20.0, ---, 40.0, ---, 30.0, ---, 20.0, 30.0, ---], + "geography.longitudes": [20.0, 40.0, 30.0, 20.0, 30.0], } @pytest.fixture def source(test_source): FIELD_SPECS = [ - {"param": "t", "step": i, "values": values.copy(), **INPUT_METADATA} for i, values in enumerate(INPUT_VALUES) + {"parameter.variable": "t", "time.step": i, "data.values": values.copy(), **INPUT_METADATA} + for i, values in enumerate(INPUT_VALUES) ] return test_source(FIELD_SPECS) @@ -71,9 +72,9 @@ def test_remove_nans_fields(source): assert np.array_equal(input_field.to_numpy(flatten=True), INPUT_VALUES[i].flatten(), equal_nan=True) assert np.array_equal(output_field.to_numpy(flatten=True), EXPECTED_VALUES[i], equal_nan=True) - output_lats, output_lons = output_field.grid_points() - assert np.array_equal(output_lats, EXPECTED_METADATA["latitudes"], equal_nan=True) - assert np.array_equal(output_lons, EXPECTED_METADATA["longitudes"], equal_nan=True) + output_lats, output_lons = output_field.geography.latlons() + assert np.array_equal(output_lats, EXPECTED_METADATA["geography.latitudes"], equal_nan=True) + assert np.array_equal(output_lons, EXPECTED_METADATA["geography.longitudes"], equal_nan=True) def test_drop_nans_tabular(): diff --git a/tests/dispatching_filters/test_rename.py b/tests/dispatching_filters/test_rename.py index a5daa16b..ed31c901 100644 --- a/tests/dispatching_filters/test_rename.py +++ b/tests/dispatching_filters/test_rename.py @@ -8,10 +8,10 @@ # nor does it submit to any jurisdiction. +import earthkit.data as ekd import pandas as pd import pytest -from anemoi.transform.fields import WrappedField from anemoi.transform.filters import create_filter_by_name as create_filter @@ -45,10 +45,10 @@ def test_rename_field(grib_source): pipeline = grib_source | rename for original, result in zip(grib_source, pipeline): - assert isinstance(result, WrappedField) - if original.metadata("param") == "z": - assert result.metadata("param") == "geopotential" - elif original.metadata("param") == "t": - assert result.metadata("param") == "temperature" + assert isinstance(result, ekd.Field) + if original.parameter.variable() == "z": + assert result.parameter.variable() == "geopotential" + elif original.parameter.variable() == "t": + assert result.parameter.variable() == "temperature" else: raise RuntimeError(f"Unexpected param: {original.metadata('param')}") diff --git a/tests/field_filters/test_accum_to_interval.py b/tests/field_filters/test_accum_to_interval.py index 5579ab84..693ca9da 100644 --- a/tests/field_filters/test_accum_to_interval.py +++ b/tests/field_filters/test_accum_to_interval.py @@ -6,6 +6,7 @@ # In applying this licence, ECMWF does not waive the privileges and immunities # granted to it by virtue of its status as an intergovernmental organisation # nor does it submit to any jurisdiction. +import datetime import numpy as np @@ -14,11 +15,15 @@ from ..utils import collect_fields_by_param MOCK_FIELD_METADATA = { - "latitudes": [10.0, 0.0, -10.0], - "longitudes": [20, 40.0], + "geography.distinct_latitudes": [10.0, 0.0, -10.0], + "geography.distinct_longitudes": [20, 40.0], } +def _to_datetime(dt_str): + return datetime.datetime.fromisoformat(dt_str.replace("Z", "+00:00")) + + def test_accum_to_interval_zero_left_true(test_source): """Accumulated-from-start fields are differenced into intervals with zero at first step. @@ -36,39 +41,34 @@ def test_accum_to_interval_zero_left_true(test_source): # Provide inputs out of chronological order to ensure sorting by valid_datetime works FIELD_SPECS = [ { - "param": "tp", - "shortName": "tp", - "values": ACC_12, - "valid_datetime": "2018-08-01T12:00:00Z", + "parameter.variable": "tp", + "data.values": ACC_12, + "time.valid_datetime": _to_datetime("2018-08-01T12:00:00Z"), **MOCK_FIELD_METADATA, }, { - "param": "tp", - "shortName": "tp", - "values": ACC_00, - "valid_datetime": "2018-08-01T00:00:00Z", + "parameter.variable": "tp", + "data.values": ACC_00, + "time.valid_datetime": _to_datetime("2018-08-01T00:00:00Z"), **MOCK_FIELD_METADATA, }, { - "param": "tp", - "shortName": "tp", - "values": ACC_06, - "valid_datetime": "2018-08-01T06:00:00Z", + "parameter.variable": "tp", + "data.values": ACC_06, + "time.valid_datetime": _to_datetime("2018-08-01T06:00:00Z"), **MOCK_FIELD_METADATA, }, # Non-target variable should pass through unchanged { - "param": "t", - "shortName": "t", - "values": B + 10, - "valid_datetime": "2018-08-01T00:00:00Z", + "parameter.variable": "t", + "data.values": B + 10, + "time.valid_datetime": _to_datetime("2018-08-01T00:00:00Z"), **MOCK_FIELD_METADATA, }, { - "param": "t", - "shortName": "t", - "values": B + 20, - "valid_datetime": "2018-08-01T06:00:00Z", + "parameter.variable": "t", + "data.values": B + 20, + "time.valid_datetime": _to_datetime("2018-08-01T06:00:00Z"), **MOCK_FIELD_METADATA, }, ] @@ -81,7 +81,7 @@ def test_accum_to_interval_zero_left_true(test_source): # Expect intervals in chronological order: 00Z -> 06Z -> 12Z assert set(output_fields) == {"tp", "t"} - tp_fields = sorted(output_fields["tp"], key=lambda f: f.metadata("valid_datetime")) + tp_fields = sorted(output_fields["tp"], key=lambda f: f.time.valid_datetime()) expected_tp = [ np.zeros_like(B), # first step zeroed @@ -91,13 +91,11 @@ def test_accum_to_interval_zero_left_true(test_source): for f, exp in zip(tp_fields, expected_tp): assert np.allclose(f.to_numpy(), exp) - # Non-target variable t should be unchanged (match by valid_datetime) - def _norm_ts(s): - return s[:-1] + "+00:00" if isinstance(s, str) and s.endswith("Z") else s - - t_inputs = {_norm_ts(spec["valid_datetime"]): spec["values"] for spec in FIELD_SPECS if spec["param"] == "t"} + t_inputs = { + spec["time.valid_datetime"]: spec["data.values"] for spec in FIELD_SPECS if spec["parameter.variable"] == "t" + } for f in output_fields["t"]: - ts = _norm_ts(f.metadata("valid_datetime")) + ts = f.time.valid_datetime() assert np.allclose(f.to_numpy(), t_inputs[ts]) @@ -112,24 +110,21 @@ def test_accum_to_interval_zero_left_false(test_source): FIELD_SPECS = [ { - "param": "tp", - "shortName": "tp", - "values": ACC_12, - "valid_datetime": "2018-08-01T12:00:00Z", + "parameter.variable": "tp", + "data.values": ACC_12, + "time.valid_datetime": _to_datetime("2018-08-01T12:00:00Z"), **MOCK_FIELD_METADATA, }, { - "param": "tp", - "shortName": "tp", - "values": ACC_06, - "valid_datetime": "2018-08-01T06:00:00Z", + "parameter.variable": "tp", + "data.values": ACC_06, + "time.valid_datetime": _to_datetime("2018-08-01T06:00:00Z"), **MOCK_FIELD_METADATA, }, { - "param": "tp", - "shortName": "tp", - "values": ACC_00, - "valid_datetime": "2018-08-01T00:00:00Z", + "parameter.variable": "tp", + "data.values": ACC_00, + "time.valid_datetime": _to_datetime("2018-08-01T00:00:00Z"), **MOCK_FIELD_METADATA, }, ] @@ -141,7 +136,7 @@ def test_accum_to_interval_zero_left_false(test_source): output_fields = collect_fields_by_param(pipeline) assert set(output_fields) == {"tp"} - tp_fields = sorted(output_fields["tp"], key=lambda f: f.metadata("valid_datetime")) + tp_fields = sorted(output_fields["tp"], key=lambda f: f.time.valid_datetime()) expected_tp = [ ACC_00, # first step unchanged when zero_left is False diff --git a/tests/field_filters/test_apply_mask.py b/tests/field_filters/test_apply_mask.py index 727375a7..cbdda88a 100644 --- a/tests/field_filters/test_apply_mask.py +++ b/tests/field_filters/test_apply_mask.py @@ -17,9 +17,9 @@ from ..utils import collect_fields_by_param MOCK_FIELD_METADATA = { - "latitudes": [10.0, 0.0, -10.0], - "longitudes": [20, 40.0], - "valid_datetime": "2018-08-01T09:00:00Z", + "geography.distinct_latitudes": [10.0, 0.0, -10.0], + "geography.distinct_longitudes": [20, 40.0], + "time.valid_datetime": "2018-08-01T09:00:00Z", } MASK_VALUES = { @@ -39,7 +39,8 @@ @pytest.fixture() def source(test_source): FIELD_SPECS = [ - {"param": param, "values": values.copy(), **MOCK_FIELD_METADATA} for param, values in DATA_VALUES.items() + {"parameter.variable": param, "data.values": values.copy(), **MOCK_FIELD_METADATA} + for param, values in DATA_VALUES.items() ] return test_source(FIELD_SPECS) @@ -53,7 +54,12 @@ def side_effect(source_type, path): # mask expected to be flattened mask = MASK_VALUES[path].copy().flatten() mock_field.to_numpy.return_value = mask - return [mock_field] + # Return a mock that supports .to_fieldlist()[0] + mock_source = mock.Mock() + mock_fieldlist = mock.Mock() + mock_fieldlist.__getitem__ = mock.Mock(return_value=mock_field) + mock_source.to_fieldlist.return_value = mock_fieldlist + return mock_source with mock.patch("anemoi.transform.filters.fields.apply_mask.ekd.from_source", autospec=True) as mock_fn: mock_fn.side_effect = side_effect diff --git a/tests/field_filters/test_apply_mask_from_field.py b/tests/field_filters/test_apply_mask_from_field.py index 9f64148d..e3de22ba 100644 --- a/tests/field_filters/test_apply_mask_from_field.py +++ b/tests/field_filters/test_apply_mask_from_field.py @@ -15,9 +15,9 @@ from ..utils import collect_fields_by_param MOCK_FIELD_METADATA = { - "latitudes": [10.0, 0.0, -10.0], - "longitudes": [20, 40.0], - "valid_datetime": "2018-08-01T09:00:00Z", + "geography.distinct_latitudes": [10.0, 0.0, -10.0], + "geography.distinct_longitudes": [20, 40.0], + "time.valid_datetime": "2018-08-01T09:00:00Z", } LSM_VALUES = np.array([[1, 0], [1, 1], [0, 0]]) @@ -32,7 +32,8 @@ @pytest.fixture() def source(test_source): FIELD_SPECS = [ - {"param": param, "values": values.copy(), **MOCK_FIELD_METADATA} for param, values in DATA_VALUES.items() + {"parameter.variable": param, "data.values": values.copy(), **MOCK_FIELD_METADATA} + for param, values in DATA_VALUES.items() ] return test_source(FIELD_SPECS) diff --git a/tests/field_filters/test_clear_step.py b/tests/field_filters/test_clear_step.py index e0c3c62c..0f1ea4fc 100644 --- a/tests/field_filters/test_clear_step.py +++ b/tests/field_filters/test_clear_step.py @@ -16,11 +16,12 @@ from anemoi.transform.filters import create_filter_by_name as create_filter from ..utils import collect_fields_by_param +from ..utils import group_component_dict MOCK_FIELD_METADATA = { - "latitudes": [10.0, 0.0, -10.0], - "longitudes": [20, 40.0], - "valid_datetime": "2018-08-01T12:00:00Z", + "geography.distinct_latitudes": [10.0, 0.0, -10.0], + "geography.distinct_longitudes": [20, 40.0], + "time.valid_datetime": "2018-08-01T12:00:00Z", } MOCK_VALUES = np.array([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]]) @@ -29,11 +30,11 @@ @pytest.fixture def source(test_source): FIELD_SPECS = [ - {"param": "t", "step": 0, "values": MOCK_VALUES, **MOCK_FIELD_METADATA}, - {"param": "t", "step": 6, "values": MOCK_VALUES, **MOCK_FIELD_METADATA}, - {"param": "t", "step": 12, "values": MOCK_VALUES, **MOCK_FIELD_METADATA}, + {"parameter.variable": "t", "time.step": 0, "data.values": MOCK_VALUES, **MOCK_FIELD_METADATA}, + {"parameter.variable": "t", "time.step": 6, "data.values": MOCK_VALUES, **MOCK_FIELD_METADATA}, + {"parameter.variable": "t", "time.step": 12, "data.values": MOCK_VALUES, **MOCK_FIELD_METADATA}, ] - return test_source(FIELD_SPECS) + return test_source([group_component_dict(s) for s in FIELD_SPECS]) def test_clear_step(source): @@ -50,13 +51,13 @@ def test_clear_step(source): assert len(output_fields[param]) == 3 for input_field, output_field in zip(input_fields[param], output_fields[param]): - input_validtime = to_datetime(input_field.metadata("valid_datetime")) - output_validtime = to_datetime(output_field.metadata("valid_datetime")) - input_step = input_field.metadata("step") + input_validtime = to_datetime(input_field.time.valid_datetime()) + output_validtime = to_datetime(output_field.time.valid_datetime()) + input_step = input_field.time.step() - expected_validtime = input_validtime - datetime.timedelta(hours=input_step) + expected_validtime = input_validtime - input_step assert output_validtime == expected_validtime - assert output_field.metadata("step") == 0 + assert output_field.time.step() == datetime.timedelta(hours=0) assert np.array_equal(input_field.to_numpy(), output_field.to_numpy()) diff --git a/tests/field_filters/test_clipper.py b/tests/field_filters/test_clipper.py index 8adeaac7..7d4b3b66 100644 --- a/tests/field_filters/test_clipper.py +++ b/tests/field_filters/test_clipper.py @@ -19,7 +19,7 @@ def calc_stats(fieldlist): stats = {} for param in ("2t", "sp"): - fields = fieldlist.sel(param=param) + fields = fieldlist.sel(**{"parameter.variable": param}) assert len(fields) == 1 data = fields[0].to_numpy() stats[param] = {"min": np.min(data), "max": np.max(data)} diff --git a/tests/field_filters/test_cos_sin_from_rad.py b/tests/field_filters/test_cos_sin_from_rad.py index 0072473c..d9f33a6f 100644 --- a/tests/field_filters/test_cos_sin_from_rad.py +++ b/tests/field_filters/test_cos_sin_from_rad.py @@ -15,9 +15,9 @@ from ..utils import collect_fields_by_param MOCK_FIELD_METADATA = { - "latitudes": [10.0, 0.0, -10.0], - "longitudes": [20, 40.0], - "valid_datetime": "2018-08-01T09:00:00Z", + "geography.distinct_latitudes": [10.0, 0.0, -10.0], + "geography.distinct_longitudes": [20, 40.0], + "time.valid_datetime": "2018-08-01T09:00:00Z", } RAD_VALUES = np.array([[2.67687254, 2.59108576], [1.83746659, 1.73104875], [1.1348185, 2.23051268]]) @@ -30,7 +30,7 @@ @pytest.fixture def RAD_source(test_source): RAD_SPEC = [ - {"param": "RAD", "values": RAD_VALUES, **MOCK_FIELD_METADATA}, + {"parameter.variable": "RAD", "data.values": RAD_VALUES, **MOCK_FIELD_METADATA}, ] return test_source(RAD_SPEC) @@ -38,7 +38,7 @@ def RAD_source(test_source): @pytest.fixture def DEG_source(test_source): DEG_SPEC = [ - {"param": "DEG", "values": np.rad2deg(RAD_VALUES), **MOCK_FIELD_METADATA}, + {"parameter.variable": "DEG", "data.values": np.rad2deg(RAD_VALUES), **MOCK_FIELD_METADATA}, ] return test_source(DEG_SPEC) @@ -46,8 +46,8 @@ def DEG_source(test_source): @pytest.fixture def cos_sin_RAD_source(test_source): COS_SIN_RAD = [ - {"param": "cos_RAD", "values": COS_RAD_VALUES, **MOCK_FIELD_METADATA}, - {"param": "sin_RAD", "values": SIN_RAD_VALUES, **MOCK_FIELD_METADATA}, + {"parameter.variable": "cos_RAD", "data.values": COS_RAD_VALUES, **MOCK_FIELD_METADATA}, + {"parameter.variable": "sin_RAD", "data.values": SIN_RAD_VALUES, **MOCK_FIELD_METADATA}, ] return test_source(COS_SIN_RAD) diff --git a/tests/field_filters/test_cos_sin_mean_wave_direction.py b/tests/field_filters/test_cos_sin_mean_wave_direction.py index b2688cc0..fb4a708a 100644 --- a/tests/field_filters/test_cos_sin_mean_wave_direction.py +++ b/tests/field_filters/test_cos_sin_mean_wave_direction.py @@ -15,9 +15,9 @@ from ..utils import collect_fields_by_param MOCK_FIELD_METADATA = { - "latitudes": [10.0, 0.0, -10.0], - "longitudes": [20, 40.0], - "valid_datetime": "2018-08-01T09:00:00Z", + "geography.distinct_latitudes": [10.0, 0.0, -10.0], + "geography.distinct_longitudes": [20, 40.0], + "time.valid_datetime": "2018-08-01T09:00:00Z", } MWD_VALUES = np.array([[153.37349864, 148.45827835], [105.27908047, 99.18178736], [65.02031089, 127.79896253]]) @@ -29,7 +29,7 @@ @pytest.fixture def mwd_source(test_source): MWD_SPEC = [ - {"param": "mwd", "values": MWD_VALUES, **MOCK_FIELD_METADATA}, + {"parameter.variable": "mwd", "data.values": MWD_VALUES, **MOCK_FIELD_METADATA}, ] return test_source(MWD_SPEC) @@ -37,8 +37,8 @@ def mwd_source(test_source): @pytest.fixture def cos_sin_mwd_source(test_source): COS_SIN_MWD = [ - {"param": "cos_mwd", "values": COS_MWD_VALUES, **MOCK_FIELD_METADATA}, - {"param": "sin_mwd", "values": SIN_MWD_VALUES, **MOCK_FIELD_METADATA}, + {"parameter.variable": "cos_mwd", "data.values": COS_MWD_VALUES, **MOCK_FIELD_METADATA}, + {"parameter.variable": "sin_mwd", "data.values": SIN_MWD_VALUES, **MOCK_FIELD_METADATA}, ] return test_source(COS_SIN_MWD) diff --git a/tests/field_filters/test_dewpoint.py b/tests/field_filters/test_dewpoint.py index 765f91d7..4ab6f2cb 100644 --- a/tests/field_filters/test_dewpoint.py +++ b/tests/field_filters/test_dewpoint.py @@ -18,9 +18,9 @@ from ..utils import collect_fields_by_param MOCK_FIELD_METADATA = { - "latitudes": [10.0, 0.0, -10.0], - "longitudes": [20, 40.0], - "valid_datetime": "2018-08-01T09:00:00Z", + "geography.distinct_latitudes": [10.0, 0.0, -10.0], + "geography.distinct_longitudes": [20, 40.0], + "time.valid_datetime": "2018-08-01T09:00:00Z", } R_VALUES = np.array([[78.13834333, 71.28598853], [99.17328572, 44.52144788], [56.49667261, 86.10495618]]) @@ -32,8 +32,8 @@ @pytest.fixture def relative_humidity_source(test_source): RELATIVE_HUMIDITY_SPEC = [ - {"param": "r", "values": R_VALUES, **MOCK_FIELD_METADATA}, - {"param": "t", "values": T_VALUES, **MOCK_FIELD_METADATA}, + {"parameter.variable": "r", "data.values": R_VALUES, **MOCK_FIELD_METADATA}, + {"parameter.variable": "t", "data.values": T_VALUES, **MOCK_FIELD_METADATA}, ] return test_source(RELATIVE_HUMIDITY_SPEC) @@ -41,8 +41,8 @@ def relative_humidity_source(test_source): @pytest.fixture def dewpoint_source(test_source): DEWPOINT_SPEC = [ - {"param": "d", "values": D_VALUES, **MOCK_FIELD_METADATA}, - {"param": "t", "values": T_VALUES, **MOCK_FIELD_METADATA}, + {"parameter.variable": "d", "data.values": D_VALUES, **MOCK_FIELD_METADATA}, + {"parameter.variable": "t", "data.values": T_VALUES, **MOCK_FIELD_METADATA}, ] return test_source(DEWPOINT_SPEC) diff --git a/tests/field_filters/test_glacier_mask.py b/tests/field_filters/test_glacier_mask.py index 6a2212b9..3b5b9445 100644 --- a/tests/field_filters/test_glacier_mask.py +++ b/tests/field_filters/test_glacier_mask.py @@ -15,11 +15,12 @@ from anemoi.transform.filters import create_filter_by_name as create_filter from ..utils import collect_fields_by_param +from ..utils import group_component_dict MOCK_FIELD_METADATA = { - "latitudes": [10.0, 0.0, -10.0], - "longitudes": [20, 40.0], - "valid_datetime": "2018-08-01T09:00:00Z", + "geography.distinct_latitudes": [10.0, 0.0, -10.0], + "geography.distinct_longitudes": [20, 40.0], + "time.valid_datetime": "2018-08-01T09:00:00Z", } SNOW_DEPTH_VALUES = np.array([[100.0, 200.0], [300.0, 400.0], [500.0, 600.0]]) @@ -28,13 +29,14 @@ @pytest.fixture def snow_depth_source(test_source): - SNOW_DEPTH_SPEC = [{"param": "sd", "values": SNOW_DEPTH_VALUES.copy(), **MOCK_FIELD_METADATA}] + SNOW_DEPTH_SPEC = [{"parameter.variable": "sd", "data.values": SNOW_DEPTH_VALUES.copy(), **MOCK_FIELD_METADATA}] return test_source(SNOW_DEPTH_SPEC) @pytest.fixture def mock_mask(): - field = {"param": "glacier_mask", "values": GLACIER_MASK.copy(), **MOCK_FIELD_METADATA} + field = {"parameter.variable": "glacier_mask", "data.values": GLACIER_MASK.copy(), **MOCK_FIELD_METADATA} + field = group_component_dict(field) return ekd.from_source("list-of-dicts", [field]) diff --git a/tests/field_filters/test_height_level_humidity.py b/tests/field_filters/test_height_level_humidity.py index 21837c50..9090417c 100644 --- a/tests/field_filters/test_height_level_humidity.py +++ b/tests/field_filters/test_height_level_humidity.py @@ -10,9 +10,9 @@ from ..utils import collect_fields_by_param MOCK_FIELD_METADATA = { - "latitudes": [10.0, 0.0, -10.0], - "longitudes": [20.0, 40.0, 60.0, 80.0], - "valid_datetime": "2018-08-01T09:00:00Z", + "geography.distinct_latitudes": [10.0, 0.0, -10.0], + "geography.distinct_longitudes": [20.0, 40.0, 60.0, 80.0], + "time.valid_datetime": "2018-08-01T09:00:00Z", } R2M_VALUES = np.array([[0, 10, 20, 30], [40, 50, 60, 70], [80, 90, 100, 110]]) @@ -75,17 +75,29 @@ @pytest.fixture def relative_humidity_source(test_source): HEIGHT_LEVEL_RELATIVE_HUMIDITY_SPEC = [ - {"param": "2r", "values": R2M_VALUES, **MOCK_FIELD_METADATA}, - {"param": "sp", "values": SP_VALUES, **MOCK_FIELD_METADATA}, - {"param": "2t", "values": T2M_VALUES, **MOCK_FIELD_METADATA}, + {"parameter.variable": "2r", "data.values": R2M_VALUES, **MOCK_FIELD_METADATA}, + {"parameter.variable": "sp", "data.values": SP_VALUES, **MOCK_FIELD_METADATA}, + {"parameter.variable": "2t", "data.values": T2M_VALUES, **MOCK_FIELD_METADATA}, ] for level, values in T_VALUES.items(): HEIGHT_LEVEL_RELATIVE_HUMIDITY_SPEC.append( - {"param": "t", "levtype": "ml", "levelist": level, "values": values, **MOCK_FIELD_METADATA} + { + "parameter.variable": "t", + "vertical.level_type": "hybrid", + "vertical.level": level, + "data.values": values, + **MOCK_FIELD_METADATA, + } ) for level, values in Q_VALUES.items(): HEIGHT_LEVEL_RELATIVE_HUMIDITY_SPEC.append( - {"param": "q", "levtype": "ml", "levelist": level, "values": values, **MOCK_FIELD_METADATA} + { + "parameter.variable": "q", + "vertical.level_type": "hybrid", + "vertical.level": level, + "data.values": values, + **MOCK_FIELD_METADATA, + } ) return test_source(HEIGHT_LEVEL_RELATIVE_HUMIDITY_SPEC) @@ -93,17 +105,29 @@ def relative_humidity_source(test_source): @pytest.fixture def specific_humidity_source(test_source): HEIGHT_LEVEL_SPECIFIC_HUMIDITY_SPEC = [ - {"param": "2sh", "values": Q2M_VALUES, **MOCK_FIELD_METADATA}, - {"param": "sp", "values": SP_VALUES, **MOCK_FIELD_METADATA}, - {"param": "2t", "values": T2M_VALUES, **MOCK_FIELD_METADATA}, + {"parameter.variable": "2sh", "data.values": Q2M_VALUES, **MOCK_FIELD_METADATA}, + {"parameter.variable": "sp", "data.values": SP_VALUES, **MOCK_FIELD_METADATA}, + {"parameter.variable": "2t", "data.values": T2M_VALUES, **MOCK_FIELD_METADATA}, ] for level, values in T_VALUES.items(): HEIGHT_LEVEL_SPECIFIC_HUMIDITY_SPEC.append( - {"param": "t", "levtype": "ml", "levelist": level, "values": values, **MOCK_FIELD_METADATA} + { + "parameter.variable": "t", + "vertical.level_type": "hybrid", + "vertical.level": level, + "data.values": values, + **MOCK_FIELD_METADATA, + } ) for level, values in Q_VALUES.items(): HEIGHT_LEVEL_SPECIFIC_HUMIDITY_SPEC.append( - {"param": "q", "levtype": "ml", "levelist": level, "values": values, **MOCK_FIELD_METADATA} + { + "parameter.variable": "q", + "vertical.level_type": "hybrid", + "vertical.level": level, + "data.values": values, + **MOCK_FIELD_METADATA, + } ) return test_source(HEIGHT_LEVEL_SPECIFIC_HUMIDITY_SPEC) @@ -111,16 +135,28 @@ def specific_humidity_source(test_source): @pytest.fixture def dewpoint_temperature_source(test_source): HEIGHT_LEVEL_DEWPOINT_TEMPERATURE_SPEC = [ - {"param": "2d", "values": D2M_VALUES, **MOCK_FIELD_METADATA}, - {"param": "sp", "values": SP_VALUES, **MOCK_FIELD_METADATA}, + {"parameter.variable": "2d", "data.values": D2M_VALUES, **MOCK_FIELD_METADATA}, + {"parameter.variable": "sp", "data.values": SP_VALUES, **MOCK_FIELD_METADATA}, ] for level, values in T_VALUES.items(): HEIGHT_LEVEL_DEWPOINT_TEMPERATURE_SPEC.append( - {"param": "t", "levtype": "ml", "levelist": level, "values": values, **MOCK_FIELD_METADATA} + { + "parameter.variable": "t", + "vertical.level_type": "hybrid", + "vertical.level": level, + "data.values": values, + **MOCK_FIELD_METADATA, + } ) for level, values in Q_VALUES.items(): HEIGHT_LEVEL_DEWPOINT_TEMPERATURE_SPEC.append( - {"param": "q", "levtype": "ml", "levelist": level, "values": values, **MOCK_FIELD_METADATA} + { + "parameter.variable": "q", + "vertical.level_type": "hybrid", + "vertical.level": level, + "data.values": values, + **MOCK_FIELD_METADATA, + } ) return test_source(HEIGHT_LEVEL_DEWPOINT_TEMPERATURE_SPEC) @@ -253,9 +289,12 @@ def test_relative_humidity_to_specific_humidity_from_file(test_source): source = test_source("anemoi-transform/filters/input_single_level_specific_humidity_to_relative_humidity.grib") input_relative_humidity = test_source("anemoi-transform/filters/single_level_relative_humidity.npy").ds.to_numpy() - md = source.ds.sel(param="2d")[0].metadata().override(edition=2, shortName="2r") + template_field = source.ds.sel(**{"parameter.variable": "2d"})[0] - source.ds += ekd.FieldList.from_array(input_relative_humidity, md) + from anemoi.transform.fields import new_field_from_numpy + + new_field = new_field_from_numpy(input_relative_humidity, template=template_field, param="2r") + source.ds = ekd.create_fieldlist(list(source.ds) + [new_field]) r_to_q_height = create_filter( "r_to_q_height", @@ -501,9 +540,13 @@ def test_dewpoint_to_specific_humidity_from_file(test_source): input_dewpoint_temperature = test_source( "anemoi-transform/filters/single_level_dewpoint_temperature.npy" ).ds.to_numpy() - md = source.ds.sel(param="2d")[0].metadata() - ds = source.ds.sel(param=["2sh", "2t", "sp", "q", "t"]) - ds += ekd.FieldList.from_array(input_dewpoint_temperature, md) + template_field = source.ds.sel(**{"parameter.variable": "2d"})[0] + ds = source.ds.sel(**{"parameter.variable": ["2sh", "2t", "sp", "q", "t"]}) + + from anemoi.transform.fields import new_field_from_numpy + + new_field = new_field_from_numpy(input_dewpoint_temperature, template=template_field, param="2d") + ds = ekd.create_fieldlist(list(ds) + [new_field]) source.ds = ds d_to_q_height = create_filter( @@ -565,7 +608,7 @@ def test_dewpoint_temperature_to_specific_humidity(dewpoint_temperature_source): # test pipeline output matches known good output result = output_fields["2sh"][0].to_numpy() expected_specific_humidity = Q2M_VALUES - np.testing.assert_allclose(result, expected_specific_humidity) + np.testing.assert_allclose(result, expected_specific_humidity, atol=1e-7) def test_height_level_dewpoint_temperature_to_specific_humidity_round_trip(dewpoint_temperature_source): diff --git a/tests/field_filters/test_impute_nans.py b/tests/field_filters/test_impute_nans.py index 81ab3b91..d164233e 100644 --- a/tests/field_filters/test_impute_nans.py +++ b/tests/field_filters/test_impute_nans.py @@ -15,9 +15,9 @@ from ..utils import collect_fields_by_param INPUT_METADATA = { - "latitudes": [10.0, 0.0, -10.0], - "longitudes": [20.0, 30.0, 40.0], - "valid_datetime": "2018-08-01T12:00:00Z", + "geography.distinct_latitudes": [10.0, 0.0, -10.0], + "geography.distinct_longitudes": [20.0, 30.0, 40.0], + "time.valid_datetime": "2018-08-01T12:00:00Z", } T_VALUES = np.array([[1.0, np.nan, 3.0], [np.nan, 5.0, 6.0], [7.0, np.nan, 9.0]]) @@ -29,9 +29,9 @@ def source(test_source): return test_source( [ - {"param": "t", "values": T_VALUES.copy(), **INPUT_METADATA}, - {"param": "q", "values": Q_VALUES.copy(), **INPUT_METADATA}, - {"param": "r", "values": R_VALUES.copy(), **INPUT_METADATA}, + {"parameter.variable": "t", "data.values": T_VALUES.copy(), **INPUT_METADATA}, + {"parameter.variable": "q", "data.values": Q_VALUES.copy(), **INPUT_METADATA}, + {"parameter.variable": "r", "data.values": R_VALUES.copy(), **INPUT_METADATA}, ] ) @@ -101,8 +101,8 @@ def test_impute_nans_grid_unchanged(source): output_fields = collect_fields_by_param(pipeline) # grid should be unchanged (unlike remove_nans which reduces the grid) - input_lats, input_lons = input_fields["t"][0].grid_points() - output_lats, output_lons = output_fields["t"][0].grid_points() + input_lats, input_lons = input_fields["t"][0].geography.latlons() + output_lats, output_lons = output_fields["t"][0].geography.latlons() assert np.array_equal(input_lats, output_lats) assert np.array_equal(input_lons, output_lons) diff --git a/tests/field_filters/test_lambda.py b/tests/field_filters/test_lambda.py index 769a9dc0..bdf33a65 100644 --- a/tests/field_filters/test_lambda.py +++ b/tests/field_filters/test_lambda.py @@ -34,7 +34,7 @@ def do_something(field: ekd.Field, a: float) -> ekd.Field: Any The modified field. """ - return field.clone(values=field.values * a) + return field.set(**{"data.values": field.values * a}) def undo_something(field: ekd.Field, a: float) -> ekd.Field: @@ -52,7 +52,7 @@ def undo_something(field: ekd.Field, a: float) -> ekd.Field: Any The modified field. """ - return field.clone(values=field.values / a) + return field.set(**{"data.values": field.values / a}) @skip_if_offline @@ -65,7 +65,7 @@ def test_earthkitfieldlambda(fieldlist: ekd.FieldList) -> None: The fieldlist to use for testing. """ - before_filter = {field.metadata("param"): field.to_numpy().copy() for field in fieldlist} + before_filter = {field.parameter.variable(): field.to_numpy().copy() for field in fieldlist} filter = create_filter( "earthkitfieldlambda", fn="tests.field_filters.test_lambda.do_something", @@ -75,10 +75,10 @@ def test_earthkitfieldlambda(fieldlist: ekd.FieldList) -> None: ) transformed = filter.forward(fieldlist) - after_forward = {field.metadata("param"): field.to_numpy().copy() for field in transformed} + after_forward = {field.parameter.variable(): field.to_numpy().copy() for field in transformed} untransformed = filter.backward(transformed) - after_backward = {field.metadata("param"): field.to_numpy().copy() for field in untransformed} + after_backward = {field.parameter.variable(): field.to_numpy().copy() for field in untransformed} for param in ("sp", "2t"): # round trip works diff --git a/tests/field_filters/test_lnsp_to_sp.py b/tests/field_filters/test_lnsp_to_sp.py index d70a7a47..54f2e0e5 100644 --- a/tests/field_filters/test_lnsp_to_sp.py +++ b/tests/field_filters/test_lnsp_to_sp.py @@ -16,9 +16,9 @@ from ..utils import collect_fields_by_param MOCK_FIELD_METADATA = { - "latitudes": [10.0, 0.0, -10.0], - "longitudes": [20, 40.0], - "valid_datetime": "2018-08-01T09:00:00Z", + "geography.distinct_latitudes": [10.0, 0.0, -10.0], + "geography.distinct_longitudes": [20, 40.0], + "time.valid_datetime": "2018-08-01T09:00:00Z", } LNSP_VALUES = np.array([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]]) @@ -28,13 +28,13 @@ @pytest.fixture def lnsp_source(test_source): - LNSP_SPEC = [{"param": "lnsp", "values": LNSP_VALUES, **MOCK_FIELD_METADATA}] + LNSP_SPEC = [{"parameter.variable": "lnsp", "data.values": LNSP_VALUES, **MOCK_FIELD_METADATA}] return test_source(LNSP_SPEC) @pytest.fixture def sp_source(test_source): - SP_SPEC = [{"param": "sp", "values": SP_VALUES, **MOCK_FIELD_METADATA}] + SP_SPEC = [{"parameter.variable": "sp", "data.values": SP_VALUES, **MOCK_FIELD_METADATA}] return test_source(SP_SPEC) diff --git a/tests/field_filters/test_orog_to_z.py b/tests/field_filters/test_orog_to_z.py index bf2a9b70..fb065b31 100644 --- a/tests/field_filters/test_orog_to_z.py +++ b/tests/field_filters/test_orog_to_z.py @@ -17,9 +17,9 @@ from ..utils import collect_fields_by_param MOCK_FIELD_METADATA = { - "latitudes": [10.0, 0.0, -10.0], - "longitudes": [20, 40.0], - "valid_datetime": "2018-08-01T09:00:00Z", + "geography.distinct_latitudes": [10.0, 0.0, -10.0], + "geography.distinct_longitudes": [20, 40.0], + "time.valid_datetime": "2018-08-01T09:00:00Z", } OROG_VALUES = np.array([[243.87788459, 1892.45371246], [427.80215359, 156.92873391], [2167.93458212, 338.15794671]]) @@ -29,13 +29,13 @@ @pytest.fixture def orog_source(test_source): - OROG_SPEC = [{"param": "orog", "values": OROG_VALUES, **MOCK_FIELD_METADATA}] + OROG_SPEC = [{"parameter.variable": "orog", "data.values": OROG_VALUES, **MOCK_FIELD_METADATA}] return test_source(OROG_SPEC) @pytest.fixture def z_source(test_source): - Z_SPEC = [{"param": "z", "values": Z_VALUES, **MOCK_FIELD_METADATA}] + Z_SPEC = [{"parameter.variable": "z", "data.values": Z_VALUES, **MOCK_FIELD_METADATA}] return test_source(Z_SPEC) diff --git a/tests/field_filters/test_pressure_level_humidity.py b/tests/field_filters/test_pressure_level_humidity.py index c29e7b24..17e613b1 100644 --- a/tests/field_filters/test_pressure_level_humidity.py +++ b/tests/field_filters/test_pressure_level_humidity.py @@ -19,9 +19,9 @@ from ..utils import collect_fields_by_param MOCK_FIELD_METADATA = { - "latitudes": [10.0, 0.0, -10.0], - "longitudes": [20, 40.0], - "valid_datetime": "2018-08-01T09:00:00Z", + "geography.distinct_latitudes": [10.0, 0.0, -10.0], + "geography.distinct_longitudes": [20, 40.0], + "time.valid_datetime": "2018-08-01T09:00:00Z", } T_VALUES = { @@ -43,10 +43,10 @@ @pytest.fixture def relative_humidity_source(test_source): PRESSURE_LEVEL_RELATIVE_HUMIDITY_SPEC = [ - {"param": "r", "levelist": 850, "values": R_VALUES[850], **MOCK_FIELD_METADATA}, - {"param": "t", "levelist": 850, "values": T_VALUES[850], **MOCK_FIELD_METADATA}, - {"param": "r", "levelist": 1000, "values": R_VALUES[1000], **MOCK_FIELD_METADATA}, - {"param": "t", "levelist": 1000, "values": T_VALUES[1000], **MOCK_FIELD_METADATA}, + {"parameter.variable": "r", "vertical.level": 850, "data.values": R_VALUES[850], **MOCK_FIELD_METADATA}, + {"parameter.variable": "t", "vertical.level": 850, "data.values": T_VALUES[850], **MOCK_FIELD_METADATA}, + {"parameter.variable": "r", "vertical.level": 1000, "data.values": R_VALUES[1000], **MOCK_FIELD_METADATA}, + {"parameter.variable": "t", "vertical.level": 1000, "data.values": T_VALUES[1000], **MOCK_FIELD_METADATA}, ] return test_source(PRESSURE_LEVEL_RELATIVE_HUMIDITY_SPEC) @@ -54,10 +54,10 @@ def relative_humidity_source(test_source): @pytest.fixture def specific_humidity_source(test_source): PRESSURE_LEVEL_SPECIFIC_HUMIDITY_SPEC = [ - {"param": "q", "levelist": 850, "values": Q_VALUES[850], **MOCK_FIELD_METADATA}, - {"param": "t", "levelist": 850, "values": T_VALUES[850], **MOCK_FIELD_METADATA}, - {"param": "q", "levelist": 1000, "values": Q_VALUES[1000], **MOCK_FIELD_METADATA}, - {"param": "t", "levelist": 1000, "values": T_VALUES[1000], **MOCK_FIELD_METADATA}, + {"parameter.variable": "q", "vertical.level": 850, "data.values": Q_VALUES[850], **MOCK_FIELD_METADATA}, + {"parameter.variable": "t", "vertical.level": 850, "data.values": T_VALUES[850], **MOCK_FIELD_METADATA}, + {"parameter.variable": "q", "vertical.level": 1000, "data.values": Q_VALUES[1000], **MOCK_FIELD_METADATA}, + {"parameter.variable": "t", "vertical.level": 1000, "data.values": T_VALUES[1000], **MOCK_FIELD_METADATA}, ] return test_source(PRESSURE_LEVEL_SPECIFIC_HUMIDITY_SPEC) @@ -79,7 +79,7 @@ def test_pressure_level_specific_humidity_to_relative_humidity(specific_humidity assert_fields_equal(input_field, output_field) # test new output matches expected values - results_by_level = {field.metadata("levelist"): field.to_numpy() for field in output_fields["r"]} + results_by_level = {field.vertical.level(): field.to_numpy() for field in output_fields["r"]} assert set(results_by_level) == {850, 1000} for level, result in results_by_level.items(): @@ -133,11 +133,13 @@ def test_pressure_level_specific_humidity_to_relative_humidity_from_file(test_so assert_fields_equal(input_field, output_field) # test pipeline output matches known good output - fields = sorted(output_fields["r"], key=lambda f: f.metadata("levelist")) + fields = sorted(output_fields["r"], key=lambda f: f.vertical.level()) fields = map(lambda f: f.to_numpy(), fields) result = np.stack(list(fields)).flatten() - expected_relative_humidity = test_source("anemoi-transform/filters/era_r.npy").ds.to_numpy().flatten() + expected_relative_humidity = ( + test_source("anemoi-transform/filters/era_r.npy").ds.to_fieldlist().to_numpy().flatten() + ) assert np.allclose(result, expected_relative_humidity) @@ -157,7 +159,7 @@ def test_pressure_level_relative_humidity_to_specific_humidity(relative_humidity assert_fields_equal(input_field, output_field) # test new output matches expected values - results_by_level = {field.metadata("levelist"): field.to_numpy() for field in output_fields["q"]} + results_by_level = {field.vertical.level(): field.to_numpy() for field in output_fields["q"]} assert set(results_by_level) == {850, 1000} for level, result in results_by_level.items(): @@ -210,7 +212,7 @@ def test_pressure_level_relative_humidity_to_specific_humidity_from_file_arome(t assert_fields_equal(input_field, output_field) # test pipeline output matches known good output - fields = sorted(output_fields["q"], key=lambda f: f.metadata("levelist")) + fields = sorted(output_fields["q"], key=lambda f: f.vertical.level()) fields = map(lambda f: f.to_numpy(), fields) result = np.stack(list(fields)) result = result.flatten() @@ -241,7 +243,7 @@ def test_pressure_level_relative_humidity_to_specific_humidity_from_file(test_so assert_fields_equal(input_field, output_field) # test pipeline output matches known good output - fields = sorted(output_fields["q"], key=lambda f: f.metadata("levelist")) + fields = sorted(output_fields["q"], key=lambda f: f.vertical.level()) fields = map(lambda f: f.to_numpy(), fields) result = np.stack(list(fields)).flatten() diff --git a/tests/field_filters/test_q_height_with_p.py b/tests/field_filters/test_q_height_with_p.py index dd0dc3a7..502c34c6 100644 --- a/tests/field_filters/test_q_height_with_p.py +++ b/tests/field_filters/test_q_height_with_p.py @@ -18,9 +18,9 @@ from ..utils import collect_fields_by_param MOCK_FIELD_METADATA = { - "latitudes": [10.0, 0.0, -10.0], - "longitudes": [20.0, 40.0], - "valid_datetime": "2018-08-01T09:00:00Z", + "geography.distinct_latitudes": [10.0, 0.0, -10.0], + "geography.distinct_longitudes": [20.0, 40.0], + "time.valid_datetime": "2018-08-01T09:00:00Z", } T_VALUES = np.array([[280.0, 290.0], [295.0, 285.0], [270.0, 300.0]]) @@ -35,9 +35,9 @@ def specific_humidity_source(test_source): return test_source( [ - {"param": "q", "values": Q_VALUES.copy(), **MOCK_FIELD_METADATA}, - {"param": "t", "values": T_VALUES.copy(), **MOCK_FIELD_METADATA}, - {"param": "pres", "values": P_VALUES.copy(), **MOCK_FIELD_METADATA}, + {"parameter.variable": "q", "data.values": Q_VALUES.copy(), **MOCK_FIELD_METADATA}, + {"parameter.variable": "t", "data.values": T_VALUES.copy(), **MOCK_FIELD_METADATA}, + {"parameter.variable": "pres", "data.values": P_VALUES.copy(), **MOCK_FIELD_METADATA}, ] ) @@ -46,9 +46,9 @@ def specific_humidity_source(test_source): def relative_humidity_source(test_source): return test_source( [ - {"param": "r", "values": R_VALUES.copy(), **MOCK_FIELD_METADATA}, - {"param": "t", "values": T_VALUES.copy(), **MOCK_FIELD_METADATA}, - {"param": "pres", "values": P_VALUES.copy(), **MOCK_FIELD_METADATA}, + {"parameter.variable": "r", "data.values": R_VALUES.copy(), **MOCK_FIELD_METADATA}, + {"parameter.variable": "t", "data.values": T_VALUES.copy(), **MOCK_FIELD_METADATA}, + {"parameter.variable": "pres", "data.values": P_VALUES.copy(), **MOCK_FIELD_METADATA}, ] ) diff --git a/tests/field_filters/test_remove_nans.py b/tests/field_filters/test_remove_nans.py index 8265215c..70eb9ccf 100644 --- a/tests/field_filters/test_remove_nans.py +++ b/tests/field_filters/test_remove_nans.py @@ -15,9 +15,9 @@ from ..utils import collect_fields_by_param INPUT_METADATA = { - "latitudes": [10.0, 0.0, -10.0], - "longitudes": [20.0, 30.0, 40.0], - "valid_datetime": "2018-08-01T12:00:00Z", + "geography.distinct_latitudes": [10.0, 0.0, -10.0], + "geography.distinct_longitudes": [20.0, 30.0, 40.0], + "time.valid_datetime": "2018-08-01T12:00:00Z", } INPUT_VALUES = [ @@ -37,17 +37,18 @@ EXPECTED_METADATA = { # take the original (flattened) versions and remove where there were NaNs in the first field - # "latitudes": [10.0, ---, 10.0, ---, 0.0, ---, -10.0, -10.0, ---], - "latitudes": [10.0, 10.0, 0.0, -10.0, -10.0], - # "longitudes": [20.0, ---, 40.0, ---, 30.0, ---, 20.0, 30.0, ---], - "longitudes": [20.0, 40.0, 30.0, 20.0, 30.0], + # "geography.latitudes": [10.0, ---, 10.0, ---, 0.0, ---, -10.0, -10.0, ---], + "geography.latitudes": [10.0, 10.0, 0.0, -10.0, -10.0], + # "geography.longitudes": [20.0, ---, 40.0, ---, 30.0, ---, 20.0, 30.0, ---], + "geography.longitudes": [20.0, 40.0, 30.0, 20.0, 30.0], } @pytest.fixture def source(test_source): FIELD_SPECS = [ - {"param": "t", "step": i, "values": values.copy(), **INPUT_METADATA} for i, values in enumerate(INPUT_VALUES) + {"parameter.variable": "t", "time.step": i, "data.values": values.copy(), **INPUT_METADATA} + for i, values in enumerate(INPUT_VALUES) ] return test_source(FIELD_SPECS) @@ -55,10 +56,11 @@ def source(test_source): @pytest.fixture def source_multiple_params(test_source): FIELD_SPECS = [ - {"param": "t", "step": i, "values": values.copy(), **INPUT_METADATA} for i, values in enumerate(INPUT_VALUES) + {"parameter.variable": "t", "time.step": i, "data.values": values.copy(), **INPUT_METADATA} + for i, values in enumerate(INPUT_VALUES) ] + [ # first step of "a" has more NaNs than "t" - {"param": "a", "step": i, "values": values.copy(), **INPUT_METADATA} + {"parameter.variable": "a", "time.step": i, "data.values": values.copy(), **INPUT_METADATA} for i, values in enumerate(INPUT_VALUES[::-1]) ] source = test_source(FIELD_SPECS) @@ -92,9 +94,9 @@ def test_remove_nans(source): assert np.array_equal(input_field.to_numpy(flatten=True), INPUT_VALUES[i].flatten(), equal_nan=True) assert np.array_equal(output_field.to_numpy(flatten=True), EXPECTED_VALUES[i], equal_nan=True) - output_lats, output_lons = output_field.grid_points() - assert np.array_equal(output_lats, EXPECTED_METADATA["latitudes"], equal_nan=True) - assert np.array_equal(output_lons, EXPECTED_METADATA["longitudes"], equal_nan=True) + output_lats, output_lons = output_field.geography.latlons() + assert np.array_equal(output_lats, EXPECTED_METADATA["geography.latitudes"], equal_nan=True) + assert np.array_equal(output_lons, EXPECTED_METADATA["geography.longitudes"], equal_nan=True) def test_remove_nans_invalid_method(): diff --git a/tests/field_filters/test_rename.py b/tests/field_filters/test_rename.py index 40ccc7a7..996d4b0b 100644 --- a/tests/field_filters/test_rename.py +++ b/tests/field_filters/test_rename.py @@ -31,12 +31,12 @@ def test_rename_grib_dict_rename(grib_source): pipeline = grib_source | rename for original, result in zip(grib_source, pipeline): - if original.metadata("param") == "z": - assert result.metadata("param") == "geopotential" - elif original.metadata("param") == "t": - assert result.metadata("param") == "temperature" + if original.parameter.variable() == "z": + assert result.parameter.variable() == "geopotential" + elif original.parameter.variable() == "t": + assert result.parameter.variable() == "temperature" else: - raise RuntimeError(f"Unexpected param: {original.metadata('param')}") + raise RuntimeError(f"Unexpected param: {original.parameter.variable()}") @skip_if_offline @@ -48,12 +48,13 @@ def test_rename_grib_format_rename(grib_source): pipeline = grib_source | rename for original, result in zip(grib_source, pipeline): - orig_param, orig_level, orig_levtype, orig_level_d = original.metadata( - "param", "levelist", "levtype", "levelist:d" - ) + orig_param = original.metadata("param") + orig_level = original.metadata("levelist") + orig_levtype = original.metadata("levtype") + orig_level_d = original.metadata("levelist:d") assert isinstance(orig_level, int) assert isinstance(orig_level_d, float) - assert result.metadata("param") == f"{orig_param}_{orig_level}_{orig_levtype}_{orig_level_d}" + assert result.parameter.variable() == f"{orig_param}_{orig_level}_{orig_levtype}_{orig_level_d}" @skip_if_offline @@ -66,13 +67,13 @@ def test_rename_grib_dict_multiple(grib_source): pipeline = grib_source | rename for original, result in zip(grib_source, pipeline): - assert result.metadata("levelist") == f"{original.metadata('levelist')}hPa" - if original.metadata("param") == "z": - assert result.metadata("param") == "geopotential" - elif original.metadata("param") == "t": - assert result.metadata("param") == "temperature" + assert result.vertical.level() == f"{original.vertical.level()}hPa" + if original.parameter.variable() == "z": + assert result.parameter.variable() == "geopotential" + elif original.parameter.variable() == "t": + assert result.parameter.variable() == "temperature" else: - raise RuntimeError(f"Unexpected param: {original.metadata('param')}") + raise RuntimeError(f"Unexpected param: {original.parameter.variable()}") @skip_if_offline @@ -84,12 +85,12 @@ def test_rename_netcdf(netcdf_source): pipeline = netcdf_source | rename for original, result in zip(netcdf_source, pipeline): - if original.metadata("param") == "t2m": - assert result.metadata("param") == "2m temperature" - elif original.metadata("param") == "msl": - assert result.metadata("param") == "mean sea level pressure" + if original.parameter.variable() == "t2m": + assert result.parameter.variable() == "2m temperature" + elif original.parameter.variable() == "msl": + assert result.parameter.variable() == "mean sea level pressure" else: - raise RuntimeError(f"Unexpected param: {original.metadata('param')}") + raise RuntimeError(f"Unexpected param: {original.parameter.variable()}") if __name__ == "__main__": diff --git a/tests/field_filters/test_repeat_members.py b/tests/field_filters/test_repeat_members.py index 94b5bbe2..e3d66936 100644 --- a/tests/field_filters/test_repeat_members.py +++ b/tests/field_filters/test_repeat_members.py @@ -16,7 +16,20 @@ from anemoi.transform.filters import create_filter_by_name as create_filter -NO_MARS = not os.path.exists(os.path.expanduser("~/.ecmwfapirc")) + +def _mars_available() -> bool: + if not os.path.exists(os.path.expanduser("~/.ecmwfapirc")): + return False + + try: + import ecmwfapi # noqa: F401 + + return True + except ImportError: + return False + + +NO_MARS = not _mars_available() def _get_template() -> tuple[Any, np.ndarray, Any]: @@ -25,11 +38,11 @@ def _get_template() -> tuple[Any, np.ndarray, Any]: Returns ------- Tuple - A tuple containing the fieldlist, values, and metadata. + A tuple containing the fieldlist, and values """ temp = ekd.from_source("mars", {"param": "2t", "levtype": "sfc", "dates": ["2023-11-17 00:00:00"]}) fieldlist = temp.to_fieldlist() - return fieldlist, fieldlist[0].values, fieldlist[0].metadata + return fieldlist, fieldlist[0].values @pytest.mark.skipif(NO_MARS, reason="No access to MARS") @@ -40,7 +53,7 @@ def test_repeat_members_using_numbers_1() -> None: - Repeating members using a list of numbers [1, 2, 3]. - Asserting the repeated members have correct values and metadata. """ - fieldlist, values, metadata = _get_template() + fieldlist, values = _get_template() repeat = create_filter("repeat_members", numbers=[1, 2, 3]) repeated = repeat.forward(fieldlist) @@ -48,8 +61,7 @@ def test_repeat_members_using_numbers_1() -> None: for i, f in enumerate(repeated): assert f.values.shape == values.shape assert np.all(f.values == values) - assert f.metadata("number") == i + 1 - assert f.metadata("name") == metadata("name") + assert f.ensemble.member() == str(i + 1) @pytest.mark.skipif(NO_MARS, reason="No access to MARS") @@ -60,7 +72,7 @@ def test_repeat_members_using_numbers_2() -> None: - Repeating members using a range of numbers "1/to/3". - Asserting the repeated members have correct values and metadata. """ - fieldlist, values, metadata = _get_template() + fieldlist, values = _get_template() repeat = create_filter("repeat_members", numbers="1/to/3") repeated = repeat.forward(fieldlist) @@ -68,8 +80,7 @@ def test_repeat_members_using_numbers_2() -> None: for i, f in enumerate(repeated): assert f.values.shape == values.shape assert np.all(f.values == values) - assert f.metadata("number") == i + 1 - assert f.metadata("name") == metadata("name") + assert f.ensemble.member() == str(i + 1) @pytest.mark.skipif(NO_MARS, reason="No access to MARS") @@ -80,7 +91,7 @@ def test_repeat_members_using_members() -> None: - Repeating members using a list of members [0, 1, 2]. - Asserting the repeated members have correct values and metadata. """ - fieldlist, values, metadata = _get_template() + fieldlist, values = _get_template() repeat = create_filter("repeat_members", members=[0, 1, 2]) repeated = repeat.forward(fieldlist) @@ -88,8 +99,7 @@ def test_repeat_members_using_members() -> None: for i, f in enumerate(repeated): assert f.values.shape == values.shape assert np.all(f.values == values) - assert f.metadata("number") == i + 1 - assert f.metadata("name") == metadata("name") + assert f.ensemble.member() == str(i + 1) @pytest.mark.skipif(NO_MARS, reason="No access to MARS") @@ -100,7 +110,7 @@ def test_repeat_members_using_count() -> None: - Repeating members using a count of 3. - Asserting the repeated members have correct values and metadata. """ - fieldlist, values, metadata = _get_template() + fieldlist, values = _get_template() repeat = create_filter("repeat_members", count=3) repeated = repeat.forward(fieldlist) @@ -108,8 +118,7 @@ def test_repeat_members_using_count() -> None: for i, f in enumerate(repeated): assert f.values.shape == values.shape assert np.all(f.values == values) - assert f.metadata("number") == i + 1 - assert f.metadata("name") == metadata("name") + assert f.ensemble.member() == str(i + 1) if __name__ == "__main__": diff --git a/tests/field_filters/test_rescale.py b/tests/field_filters/test_rescale.py index e273d1bd..4517f99c 100644 --- a/tests/field_filters/test_rescale.py +++ b/tests/field_filters/test_rescale.py @@ -24,16 +24,16 @@ def test_rescale(fieldlist: ekd.FieldList) -> None: The fieldlist to use for testing. """ - before_filter = {field.metadata("param"): field.to_numpy().copy() for field in fieldlist} + before_filter = {field.parameter.variable(): field.to_numpy().copy() for field in fieldlist} # rescale from K to °C k_to_deg = create_filter("rescale", scale=1.0, offset=-273.15, param="2t") rescaled = k_to_deg.forward(fieldlist) - after_forward = {field.metadata("param"): field.to_numpy().copy() for field in rescaled} + after_forward = {field.parameter.variable(): field.to_numpy().copy() for field in rescaled} # and back rescaled_back = k_to_deg.backward(rescaled) - after_backward = {field.metadata("param"): field.to_numpy().copy() for field in rescaled_back} + after_backward = {field.parameter.variable(): field.to_numpy().copy() for field in rescaled_back} for param in ("2t", "sp"): npt.assert_allclose(before_filter[param], after_backward[param]) @@ -53,15 +53,15 @@ def test_convert(fieldlist: ekd.FieldList) -> None: fieldlist : ekd.FieldList The fieldlist to use for testing. """ - before_filter = {field.metadata("param"): field.to_numpy().copy() for field in fieldlist} + before_filter = {field.parameter.variable(): field.to_numpy().copy() for field in fieldlist} # rescale from K to °C k_to_deg = create_filter("convert", unit_in="K", unit_out="degC", param="2t") rescaled = k_to_deg.forward(fieldlist) - after_forward = {field.metadata("param"): field.to_numpy().copy() for field in rescaled} + after_forward = {field.parameter.variable(): field.to_numpy().copy() for field in rescaled} # and back rescaled_back = k_to_deg.backward(rescaled) - after_backward = {field.metadata("param"): field.to_numpy().copy() for field in rescaled_back} + after_backward = {field.parameter.variable(): field.to_numpy().copy() for field in rescaled_back} for param in ("2t", "sp"): npt.assert_allclose(before_filter[param], after_backward[param]) diff --git a/tests/field_filters/test_rodeo_opera_clipping.py b/tests/field_filters/test_rodeo_opera_clipping.py index de1f7148..c9045b51 100644 --- a/tests/field_filters/test_rodeo_opera_clipping.py +++ b/tests/field_filters/test_rodeo_opera_clipping.py @@ -17,9 +17,9 @@ MAX_TP = 12.5 MOCK_FIELD_METADATA = { - "latitudes": [10.0, 0.0, -10.0], - "longitudes": [20, 40.0], - "valid_datetime": "2018-08-01T09:00:00Z", + "geography.distinct_latitudes": [10.0, 0.0, -10.0], + "geography.distinct_longitudes": [20, 40.0], + "time.valid_datetime": "2018-08-01T09:00:00Z", } expected_tp_values = np.array( @@ -57,8 +57,8 @@ def rodeo_opera_source(test_source): ) SPEC = [ - {"param": "tp", "values": tp_values, **MOCK_FIELD_METADATA}, - {"param": "qi", "values": qi_values, **MOCK_FIELD_METADATA}, + {"parameter.variable": "tp", "data.values": tp_values, **MOCK_FIELD_METADATA}, + {"parameter.variable": "qi", "data.values": qi_values, **MOCK_FIELD_METADATA}, ] return test_source(SPEC) diff --git a/tests/field_filters/test_rodeo_opera_preprocessing.py b/tests/field_filters/test_rodeo_opera_preprocessing.py index 9477da6a..1cdd3bbc 100644 --- a/tests/field_filters/test_rodeo_opera_preprocessing.py +++ b/tests/field_filters/test_rodeo_opera_preprocessing.py @@ -20,9 +20,9 @@ MAX_TP = 12.5 MOCK_FIELD_METADATA = { - "latitudes": [10.0, 0.0, -10.0], - "longitudes": [20, 40.0], - "valid_datetime": "2018-08-01T09:00:00Z", + "geography.distinct_latitudes": [10.0, 0.0, -10.0], + "geography.distinct_longitudes": [20, 40.0], + "time.valid_datetime": "2018-08-01T09:00:00Z", } expected_tp_values = np.array( @@ -67,9 +67,9 @@ def rodeo_opera_source(test_source): ) SPEC = [ - {"param": "tp", "values": tp_values, **MOCK_FIELD_METADATA}, - {"param": "qi", "values": qi_values, **MOCK_FIELD_METADATA}, - {"param": "dm", "values": dm_values, **MOCK_FIELD_METADATA}, + {"parameter.variable": "tp", "data.values": tp_values, **MOCK_FIELD_METADATA}, + {"parameter.variable": "qi", "data.values": qi_values, **MOCK_FIELD_METADATA}, + {"parameter.variable": "dm", "data.values": dm_values, **MOCK_FIELD_METADATA}, ] return test_source(SPEC) diff --git a/tests/field_filters/test_rotate_winds.py b/tests/field_filters/test_rotate_winds.py index e3bf8c59..dcee5e19 100644 --- a/tests/field_filters/test_rotate_winds.py +++ b/tests/field_filters/test_rotate_winds.py @@ -15,9 +15,10 @@ from ..utils import collect_fields_by_param MOCK_FIELD_METADATA = { - "latitudes": [10.0, 0.0, -10.0], - "longitudes": [20, 40.0], - "valid_datetime": "2018-08-01T09:00:00Z", + "geography.distinct_latitudes": [10.0, 0.0, -10.0], + "geography.distinct_longitudes": [20, 40.0], + "geography.projTargetString": "+proj=eqc +ellps=WGS84 +a=6378137.0 +lon_0=0.0 +to_meter=111319.4907932736 +no_defs +type=crs", + "time.valid_datetime": "2018-08-01T09:00:00Z", } U_VALUES = np.array([[-3.26786804, -2.90458679], [-4.28153992, -10.75224304], [-6.29130554, -4.17704773]]) @@ -30,8 +31,8 @@ @pytest.fixture def wind_source(test_source): WIND_SPEC = [ - {"param": "10u", "values": U_VALUES, **MOCK_FIELD_METADATA}, - {"param": "10v", "values": V_VALUES, **MOCK_FIELD_METADATA}, + {"parameter.variable": "10u", "data.values": U_VALUES, **MOCK_FIELD_METADATA}, + {"parameter.variable": "10v", "data.values": V_VALUES, **MOCK_FIELD_METADATA}, ] return test_source(WIND_SPEC) @@ -39,8 +40,8 @@ def wind_source(test_source): @pytest.fixture def rotated_wind_source(test_source): ROTATED_WIND_SPEC = [ - {"param": "10u", "values": ROTATED_U_VALUES, **MOCK_FIELD_METADATA}, - {"param": "10v", "values": ROTATED_V_VALUES, **MOCK_FIELD_METADATA}, + {"parameter.variable": "10u", "data.values": ROTATED_U_VALUES, **MOCK_FIELD_METADATA}, + {"parameter.variable": "10v", "data.values": ROTATED_V_VALUES, **MOCK_FIELD_METADATA}, ] return test_source(ROTATED_WIND_SPEC) diff --git a/tests/field_filters/test_sum.py b/tests/field_filters/test_sum.py index 2f6d853d..27a420ab 100644 --- a/tests/field_filters/test_sum.py +++ b/tests/field_filters/test_sum.py @@ -15,9 +15,9 @@ from ..utils import collect_fields_by_param MOCK_FIELD_METADATA = { - "latitudes": [10.0, 0.0, -10.0], - "longitudes": [20, 40.0], - "valid_datetime": "2018-08-01T09:00:00Z", + "geography.distinct_latitudes": [10.0, 0.0, -10.0], + "geography.distinct_longitudes": [20, 40.0], + "time.valid_datetime": "2018-08-01T09:00:00Z", } T_VALUES = np.array([[293.32301331, 284.21559143], [260.53981018, 291.18824768], [279.88941956, 248.87574768]]) @@ -26,7 +26,7 @@ R_VALUES = np.array([[37.91091442, 79.51638317], [95.61794567, 71.53396130], [70.03982067, 89.69021130]]) -EXPECTED_SUM = (R_VALUES + T_VALUES).flatten() +EXPECTED_SUM = R_VALUES + T_VALUES EXPECTED_SUM_MULTILEVEL = (T_VALUES * 2.0 - 15.0).flatten() @@ -34,9 +34,9 @@ @pytest.fixture def sum_input_source_one_level(mars_test_source): PRESSURE_LEVEL_RELATIVE_HUMIDITY_SPEC = [ - {"param": "r", "levelist": 850, "values": R_VALUES, **MOCK_FIELD_METADATA}, - {"param": "t", "levelist": 850, "values": T_VALUES, **MOCK_FIELD_METADATA}, - {"param": "q", "levelist": 850, "values": Q_VALUES, **MOCK_FIELD_METADATA}, + {"parameter.variable": "r", "vertical.level": 850, "data.values": R_VALUES, **MOCK_FIELD_METADATA}, + {"parameter.variable": "t", "vertical.level": 850, "data.values": T_VALUES, **MOCK_FIELD_METADATA}, + {"parameter.variable": "q", "vertical.level": 850, "data.values": Q_VALUES, **MOCK_FIELD_METADATA}, ] return mars_test_source(PRESSURE_LEVEL_RELATIVE_HUMIDITY_SPEC) @@ -44,9 +44,9 @@ def sum_input_source_one_level(mars_test_source): @pytest.fixture def sum_input_source_multilevel(mars_test_source): MULTILEVEL_TEMP_RELATIVE_HUMIDITY = [ - {"param": "t_850", "levelist": 850, "values": T_VALUES, **MOCK_FIELD_METADATA}, - {"param": "t_500", "levelist": 500, "values": T_VALUES - 15.0, **MOCK_FIELD_METADATA}, - {"param": "r", "levelist": 850, "values": R_VALUES, **MOCK_FIELD_METADATA}, + {"parameter.variable": "t_850", "vertical.level": 850, "data.values": T_VALUES, **MOCK_FIELD_METADATA}, + {"parameter.variable": "t_500", "vertical.level": 500, "data.values": T_VALUES - 15.0, **MOCK_FIELD_METADATA}, + {"parameter.variable": "r", "vertical.level": 850, "data.values": R_VALUES, **MOCK_FIELD_METADATA}, ] return mars_test_source(MULTILEVEL_TEMP_RELATIVE_HUMIDITY) @@ -95,8 +95,8 @@ def test_sum_fields_multilevel(sum_input_source_multilevel): # Validate the sum field # arrays are flattened in sum - assert output_fields["sum"][0].to_numpy().shape == EXPECTED_SUM_MULTILEVEL.shape - assert np.allclose(output_fields["sum"][0].to_numpy(), EXPECTED_SUM_MULTILEVEL) + assert output_fields["sum"][0].to_numpy(flatten=True).shape == EXPECTED_SUM_MULTILEVEL.shape + assert np.allclose(output_fields["sum"][0].to_numpy(flatten=True), EXPECTED_SUM_MULTILEVEL) def test_sum_multilevel_ignore_level_false(sum_input_source_multilevel): diff --git a/tests/field_filters/test_uv_to_ddff.py b/tests/field_filters/test_uv_to_ddff.py index aab62665..d359f8ed 100644 --- a/tests/field_filters/test_uv_to_ddff.py +++ b/tests/field_filters/test_uv_to_ddff.py @@ -16,9 +16,9 @@ from ..utils import collect_fields_by_param MOCK_FIELD_METADATA = { - "latitudes": [10.0, 0.0, -10.0], - "longitudes": [20, 40.0], - "valid_datetime": "2018-08-01T09:00:00Z", + "geography.distinct_latitudes": [10.0, 0.0, -10.0], + "geography.distinct_longitudes": [20, 40.0], + "time.valid_datetime": "2018-08-01T09:00:00Z", } U_VALUES = { @@ -45,10 +45,10 @@ @pytest.fixture def uv_source(test_source): UV_SPEC = [ - {"param": "u", "levelist": 500, "values": U_VALUES[500], **MOCK_FIELD_METADATA}, - {"param": "v", "levelist": 500, "values": V_VALUES[500], **MOCK_FIELD_METADATA}, - {"param": "u", "levelist": 850, "values": U_VALUES[850], **MOCK_FIELD_METADATA}, - {"param": "v", "levelist": 850, "values": V_VALUES[850], **MOCK_FIELD_METADATA}, + {"parameter.variable": "u", "vertical.level": 500, "data.values": U_VALUES[500], **MOCK_FIELD_METADATA}, + {"parameter.variable": "v", "vertical.level": 500, "data.values": V_VALUES[500], **MOCK_FIELD_METADATA}, + {"parameter.variable": "u", "vertical.level": 850, "data.values": U_VALUES[850], **MOCK_FIELD_METADATA}, + {"parameter.variable": "v", "vertical.level": 850, "data.values": V_VALUES[850], **MOCK_FIELD_METADATA}, ] return test_source(UV_SPEC) @@ -56,10 +56,10 @@ def uv_source(test_source): @pytest.fixture def ddff_source(test_source): DDFF_SPEC = [ - {"param": "ws", "levelist": 500, "values": DD_VALUES[500], **MOCK_FIELD_METADATA}, - {"param": "wdir", "levelist": 500, "values": FF_VALUES[500], **MOCK_FIELD_METADATA}, - {"param": "ws", "levelist": 850, "values": DD_VALUES[850], **MOCK_FIELD_METADATA}, - {"param": "wdir", "levelist": 850, "values": FF_VALUES[850], **MOCK_FIELD_METADATA}, + {"parameter.variable": "ws", "vertical.level": 500, "data.values": DD_VALUES[500], **MOCK_FIELD_METADATA}, + {"parameter.variable": "wdir", "vertical.level": 500, "data.values": FF_VALUES[500], **MOCK_FIELD_METADATA}, + {"parameter.variable": "ws", "vertical.level": 850, "data.values": DD_VALUES[850], **MOCK_FIELD_METADATA}, + {"parameter.variable": "wdir", "vertical.level": 850, "data.values": FF_VALUES[850], **MOCK_FIELD_METADATA}, ] return test_source(DDFF_SPEC) @@ -75,8 +75,8 @@ def test_uv_to_ddff(uv_source): for i, level in enumerate([500, 850]): assert np.allclose(output_fields["ws"][i].to_numpy(), DD_VALUES[level]) assert np.allclose(output_fields["wdir"][i].to_numpy(), FF_VALUES[level]) - assert output_fields["ws"][i].metadata("levelist") == level - assert output_fields["wdir"][i].metadata("levelist") == level + assert output_fields["ws"][i].vertical.level() == level + assert output_fields["wdir"][i].vertical.level() == level def test_uv_to_ddff_round_trip(uv_source): @@ -109,8 +109,8 @@ def test_ddff_to_uv(ddff_source): for i, level in enumerate([500, 850]): assert np.allclose(output_fields["u"][i].to_numpy(), U_VALUES[level]) assert np.allclose(output_fields["v"][i].to_numpy(), V_VALUES[level]) - assert output_fields["u"][i].metadata("levelist") == level - assert output_fields["v"][i].metadata("levelist") == level + assert output_fields["u"][i].vertical.level() == level + assert output_fields["v"][i].vertical.level() == level def test_ddff_to_uv_round_trip(ddff_source): diff --git a/tests/test_dispatchingfilter.py b/tests/test_dispatchingfilter.py index e3d9d87a..54e84056 100644 --- a/tests/test_dispatchingfilter.py +++ b/tests/test_dispatchingfilter.py @@ -6,7 +6,7 @@ TEST_CASES = [ pytest.param(pd.DataFrame(), id="dataframe"), - pytest.param(ekd.FieldList(), id="fieldlist"), + pytest.param(ekd.create_fieldlist(), id="fieldlist"), pytest.param(None, id="other"), ] diff --git a/tests/test_fields.py b/tests/test_fields.py index 7763662c..587b9a65 100644 --- a/tests/test_fields.py +++ b/tests/test_fields.py @@ -18,56 +18,60 @@ @pytest.fixture def sample_field(): - return ekd.from_source("sample", "test.grib")[0] + return ekd.from_source("sample", "test.grib").to_fieldlist()[0] +@pytest.mark.xfail(reason="setting arbitrary metadata not yet supported") def test_field_new_metadata(sample_field): """Test that a new field can be created with new metadata.""" - assert "foo" not in sample_field.metadata() + # TODO: consider whether new_field_with_metadata should allow setting ekd field labels new_field = new_field_with_metadata(sample_field, foo="bar") - assert new_field.metadata("foo") == "bar" + assert new_field.get("labels.foo") == "bar" def test_field_update_metadata(sample_field): """Test that a new field can be created with updated metadata.""" - assert sample_field.metadata("param") == "2t" + assert sample_field.parameter.variable() == "2t" new_field = new_field_with_metadata(sample_field, param="foo") - assert new_field.metadata("param") == "foo" + assert new_field.parameter.variable() == "foo" +@pytest.mark.xfail(reason="centre key not currently supported") def test_update_multiple_metadata(sample_field): """Test that we can update multiple metadata keys at once.""" - assert sample_field.metadata("param", "centre") == ("2t", "ecmf") + assert sample_field.metadata("param") == "2t" + assert sample_field.metadata("centre") == "ecmf" new_field = new_field_with_metadata(sample_field, param="foo", centre="bar") - assert new_field.metadata("param", "centre") == ("foo", "bar") + assert new_field.parameter.variable() == "foo" + assert new_field.metadata("centre") == "bar" -@pytest.mark.xfail(reason="__contains__ not yet implemented for metadata") +@pytest.mark.xfail(reason="setting arbitrary metadata not yet supported") def test_metadata_in_new_field(sample_field): """Test that we can check if a key is in the metadata.""" - assert "foo" not in sample_field.metadata() new_field = new_field_with_metadata(sample_field, foo="bar") - assert "foo" in new_field.metadata() + assert new_field.get("labels.foo") == "bar" +@pytest.mark.xfail(reason="no dict-like metadata interface") def test_field_with_updated_metadata_has_same_keys(sample_field): """Test that updating existing metadata leaves the keys unchanged.""" - assert "param" in sample_field.metadata() + assert sample_field.metadata("param") == "2t" new_field = new_field_with_metadata(sample_field, param="foo") - assert tuple(new_field.metadata().keys()) == tuple(sample_field.metadata().keys()) + assert new_field.parameter.variable() == "foo" + assert new_field.metadata("param") == "foo" -@pytest.mark.xfail(reason="updating metadata keys not yet implemented") +@pytest.mark.xfail(reason="setting arbitrary metadata not yet supported") def test_field_adding_metadata_updates_keys(sample_field): """Test that adding a new metadata key is reflected in the keys.""" - assert "foo" not in tuple(sample_field.metadata().keys()) new_field = new_field_with_metadata(sample_field, foo="bar") - assert "foo" in tuple(new_field.metadata().keys()) + assert new_field.get("labels.foo") == "bar" def test_fieldselection_match_all(): """Test FieldSelection with no arguments matches all fields.""" - field = mock_field(invalid_key="any_value") + field = mock_field(**{"labels.invalid_key": "value"}) selection = FieldSelection() assert selection.match(field) @@ -80,55 +84,55 @@ def test_fieldselection_invalid_key(): def test_fieldselection_match_fail_different_param(): """Test FieldSelection match fails when param is different.""" - field = mock_field(param="2t") - selection = FieldSelection(param="2z") + field = mock_field(**{"parameter.variable": "2t"}) + selection = FieldSelection(**{"parameter.variable": "2z"}) assert not selection.match(field) def test_fieldselection_match_same_param(): """Test FieldSelection match succeeds when param is the same.""" - field = mock_field(param="2t") - selection = FieldSelection(param="2t") + field = mock_field(**{"parameter.variable": "2t"}) + selection = FieldSelection(**{"parameter.variable": "2t"}) assert selection.match(field) def test_fieldselection_match_fail_missing_key(): """Test FieldSelection match fails when a selection key is missing on the field.""" - field = mock_field(param="t") - selection = FieldSelection(param="t", levelist=850) + field = mock_field(**{"parameter.variable": "t"}) + selection = FieldSelection(**{"parameter.variable": "2t", "vertical.level": 850}) assert not selection.match(field) def test_fieldselection_match_field_with_extra_metadata(): """Test FieldSelection match succeeds when the field has extra metadata.""" - field = mock_field(param="t", levelist=850) - selection = FieldSelection(param="t") + field = mock_field(**{"parameter.variable": "t", "vertical.level": 850}) + selection = FieldSelection(**{"parameter.variable": "t"}) assert selection.match(field) def test_fieldselection_match_fail_same_param_different_level(): """Test FieldSelection match fails when param is the same but the levelist is different.""" - field = mock_field(param="t", levelist=100) - selection = FieldSelection(param="t", levelist=850) + field = mock_field(**{"parameter.variable": "t", "vertical.level": 100}) + selection = FieldSelection(**{"parameter.variable": "t", "vertical.level": 850}) assert not selection.match(field) def test_fieldselection_match_same_param_same_level(): """Test FieldSelection match succeeds when param and level are the same.""" - field = mock_field(param="t", levelist=850) - selection = FieldSelection(param="t", levelist=850) + field = mock_field(**{"parameter.variable": "t", "vertical.level": 850}) + selection = FieldSelection(**{"parameter.variable": "t", "vertical.level": 850}) assert selection.match(field) def test_fieldselection_match_is_subset(): """Test FieldSelection match succeeds when the field is a subset of the selection.""" - field = mock_field(param="t", levelist=850) - selection = FieldSelection(param=["t", "q"], levelist=[850, 950]) + field = mock_field(**{"parameter.variable": "t", "vertical.level": 850}) + selection = FieldSelection(**{"parameter.variable": ["t", "q"], "vertical.level": [850, 950]}) assert selection.match(field) def test_fieldselection_match_fail_different_param_same_level(): """Test FieldSelection match fails when the is on the same level but a different param.""" - field = mock_field(param="t", levelist=850) - selection = FieldSelection(param="q", levelist=[850, 950]) + field = mock_field(**{"parameter.variable": "t", "vertical.level": 850}) + selection = FieldSelection(**{"parameter.variable": "q", "vertical.level": [850, 950]}) assert not selection.match(field) diff --git a/tests/test_filter.py b/tests/test_filter.py index 5b0fce02..efce4bec 100644 --- a/tests/test_filter.py +++ b/tests/test_filter.py @@ -171,17 +171,17 @@ def forward_transform(self, field): return self.new_field_from_numpy(field.to_numpy() + 1, template=field) def forward_select(self): - return {"param": self.temperature} + return {"parameter.variable": self.temperature} # source dataset has 2t and 2r variables pipeline = source | TestFilter(temperature="2t") result_params = [] for original, result in zip(source, pipeline): - result_params.append(result.metadata("param")) - assert original.metadata("param") == result.metadata("param") + result_params.append(result.parameter.variable()) + assert original.parameter.variable() == result.parameter.variable() # only 2t has transform applied - if result.metadata("param") == "2t": + if result.parameter.variable() == "2t": assert np.allclose(original.to_numpy() + 1, result.to_numpy()) else: assert np.allclose(original.to_numpy(), result.to_numpy()) @@ -202,17 +202,17 @@ def backward_transform(self, field): return self.new_field_from_numpy(field.to_numpy() - 1, template=field) def backward_select(self): - return {"param": self.temperature} + return {"parameter.variable": self.temperature} # source dataset has 2t and 2r variables pipeline = source | TestFilter.reversed(temperature="2t") result_params = [] for original, result in zip(source, pipeline): - result_params.append(result.metadata("param")) - assert original.metadata("param") == result.metadata("param") + result_params.append(result.parameter.variable()) + assert original.parameter.variable() == result.parameter.variable() # only 2t has transform applied - if result.metadata("param") == "2t": + if result.parameter.variable() == "2t": assert np.allclose(original.to_numpy() - 1, result.to_numpy()) else: assert np.allclose(original.to_numpy(), result.to_numpy()) @@ -233,16 +233,16 @@ def backward_transform(self, field): return self.new_field_from_numpy(field.to_numpy() - 1, template=field) def forward_select(self): - return {"param": self.temperature} + return {"parameter.variable": self.temperature} # source dataset has 2t and 2r variables pipeline = source | TestFilter.reversed(temperature="2t") result_params = [] for original, result in zip(source, pipeline): - result_params.append(result.metadata("param")) - assert original.metadata("param") == result.metadata("param") - if result.metadata("param") == "2t": + result_params.append(result.parameter.variable()) + assert original.parameter.variable() == result.parameter.variable() + if result.parameter.variable() == "2t": assert np.allclose(original.to_numpy() - 1, result.to_numpy()) else: assert np.allclose(original.to_numpy(), result.to_numpy()) @@ -257,10 +257,10 @@ class TestFilter(SingleFieldFilter): required_inputs = ("temperature", "renamed_temperature") def forward_select(self): - return {"param": self.temperature} + return {"parameter.variable": self.temperature} def backward_select(self): - return {"param": self.renamed_temperature} + return {"parameter.variable": self.renamed_temperature} def forward_transform(self, field): new_metadata = {"param": self.renamed_temperature} @@ -276,15 +276,15 @@ def backward_transform(self, field): pipeline = source | forward_filter # forward transform for original, result in zip(source, pipeline): - if original.metadata("param") == "2t": + if original.parameter.variable() == "2t": assert np.allclose(original.to_numpy() + 1, result.to_numpy()) - assert result.metadata("param") == "2t_renamed" + assert result.parameter.variable() == "2t_renamed" else: assert np.allclose(original.to_numpy(), result.to_numpy()) - assert original.metadata("param") == result.metadata("param") + assert original.parameter.variable() == result.parameter.variable() # round trip pipeline = source | forward_filter | backward_filter for original, result in zip(source, pipeline): assert np.allclose(original.to_numpy(), result.to_numpy()) - assert original.metadata("param") == result.metadata("param") + assert original.parameter.variable() == result.parameter.variable() diff --git a/tests/test_grids.py b/tests/test_grids.py index e0769cd1..176c9826 100644 --- a/tests/test_grids.py +++ b/tests/test_grids.py @@ -33,7 +33,7 @@ def do_not_test_unstructured_from_url() -> None: assert len(ds) == 1 - lats, lons = ds[0].grid_points() + lats, lons = ds[0].geography.latlons() assert len(lats) == len(lons) diff --git a/tests/test_grouping.py b/tests/test_grouping.py index dabf7bf3..29b919be 100644 --- a/tests/test_grouping.py +++ b/tests/test_grouping.py @@ -7,17 +7,10 @@ def field_generator(**metadata_values): - MOCK_MARS_METADATA = { - "domain": "g", - "levtype": "sfc", - "date": 20200513, - "time": 1200, - "step": 0, - "param": "2t", - "class": "od", - "type": "an", - "stream": "oper", - "expver": "0001", + MOCK_METADATA = { + "time.step": 0, + "time.valid_datetime": "2020-05-13T12:00:00Z", + "parameter.variable": "2t", } # builds fields with metadata from cartesian product of metadata_values fields = [] @@ -25,7 +18,7 @@ def field_generator(**metadata_values): combinations = itertools.product(*metadata_values.values()) for values in combinations: - metadata = MOCK_MARS_METADATA | dict(zip(metadata_values.keys(), values)) + metadata = MOCK_METADATA | dict(zip(metadata_values.keys(), values)) fields.append(mock_field(**metadata)) return fields @@ -33,15 +26,30 @@ def field_generator(**metadata_values): @pytest.fixture def sample_fields(): return field_generator( - step=[0, 1], - param=["t", "q", "u", "v"], + **{ + "time.step": [0, 1], + "parameter.variable": ["t", "q", "u", "v"], + } ) @pytest.fixture def sample_fields_vertical(): - surface_fields = field_generator(step=[0, 1], levtype=["sfc"], param=["2q", "2r", "2t", "sp"]) - vertical_fields = field_generator(step=[0, 1], levtype=["ml"], param=["q", "t"], levelist=[1, 2, 3]) + surface_fields = field_generator( + **{ + "time.step": [0, 1], + "vertical.level_type": ["sfc"], + "parameter.variable": ["2q", "2r", "2t", "sp"], + } + ) + vertical_fields = field_generator( + **{ + "time.step": [0, 1], + "vertical.level_type": ["hybrid"], + "parameter.variable": ["q", "t"], + "vertical.level": [1, 2, 3], + } + ) return surface_fields + vertical_fields @@ -55,27 +63,19 @@ def test_group_by_param(sample_fields): for group in grouper.iterate(sample_fields, other=other.append): assert len(group) == len(match_params) # ensure order is the same - assert [field.metadata("param") for field in group] == match_params - metadata = [] + assert [field.parameter.variable() for field in group] == match_params for field in group: num_matching += 1 # check field is unchanged assert field in sample_fields - # get metadata except param from each field - m = field.metadata(namespace="mars") - m.pop("param", None) - metadata.append(m) - # rest of the metadata the same within a group - assert all(m == metadata[0] for m in metadata[1:]) - assert num_matching + len(other) == len(sample_fields) for field in other: - assert field.metadata("param") not in match_params + assert field.parameter.variable() not in match_params assert field in sample_fields -@pytest.mark.xfail(reason="vertical grouping not yet implemented") +@pytest.mark.xfail(reason="vertical grouping test to be revisited") def test_group_by_param_vertical(sample_fields_vertical): from anemoi.transform.grouping import GroupByParamVertical @@ -83,7 +83,7 @@ def get_param(f): if isinstance(f, ekd.Field): f = [f] - param = [x.metadata("param") for x in f] + param = [x.parameter.variable() for x in f] assert len(set(param)) == 1 return param[0] @@ -105,11 +105,11 @@ def get_param(f): num_matching += 1 # check field is unchanged assert field in sample_fields_vertical - # get metadata except keys known to be different from each field - m = field.metadata(namespace="mars") - m.pop("param", None) - m.pop("levtype", None) - m.pop("levelist", None) + # get metadata via component API (namespace="mars" removed in ekd 1.0) + m = { + "step": field.time.step(), + "valid_datetime": field.time.valid_datetime(), + } metadata.append(m) # rest of the metadata the same within a group @@ -117,5 +117,5 @@ def get_param(f): assert num_matching + len(other) == len(sample_fields_vertical) for field in other: - assert field.metadata("param") not in match_params + assert field.parameter.variable() not in match_params assert field in sample_fields_vertical diff --git a/tests/test_matching.py b/tests/test_matching.py index 7cfe1d2c..4e206c44 100644 --- a/tests/test_matching.py +++ b/tests/test_matching.py @@ -13,6 +13,7 @@ import earthkit.data as ekd import pytest +from anemoi.transform.fields import new_field_from_numpy from anemoi.transform.filters.fields.matching import MatchingFieldsFilter from anemoi.transform.filters.fields.matching import MatchingSpec @@ -33,11 +34,7 @@ def __init__(self, *, a, b, return_inputs="none"): def forward_transform(self, a: ekd.Field, b: ekd.Field) -> Iterator[ekd.Field]: result = a.values + b.values - yield self.new_field_from_numpy(result, template=a, param="c") - - def new_field_from_numpy(self, array, *, template, param): - metadata = dict(template.metadata()) | {"param": param} - return mock_field(**metadata) + yield new_field_from_numpy(result, template=a, param="c") def test_matching_spec_initializes_correctly(): @@ -47,43 +44,53 @@ def test_matching_spec_initializes_correctly(): def test_forward_transform_adds_fields(): - a = mock_field(param="a", step=0, level=850) - b = mock_field(param="b", step=0, level=850) - data = ekd.SimpleFieldList([a, b]) + a = mock_field( + **{"parameter.variable": "a", "time.valid_datetime": "2000-01-01T00:00Z", "time.step": 0, "vertical.level": 850} + ) + b = mock_field( + **{"parameter.variable": "b", "time.valid_datetime": "2000-01-01T00:00Z", "time.step": 0, "vertical.level": 850} + ) + data = ekd.create_fieldlist([a, b]) f = AddFields(a="a", b="b") result = f.forward(data) assert len(result) == 1 assert isinstance(result[0], ekd.Field) - assert result[0].metadata("param") == "c" + assert result[0].parameter.variable() == "c" def test_return_inputs(): - a = mock_field(param="a", step=0, level=850) - b = mock_field(param="b", step=0, level=850) - data = ekd.SimpleFieldList([a, b]) + a = mock_field( + **{"parameter.variable": "a", "time.valid_datetime": "2000-01-01T00:00Z", "time.step": 0, "vertical.level": 850} + ) + b = mock_field( + **{"parameter.variable": "b", "time.valid_datetime": "2000-01-01T00:00Z", "time.step": 0, "vertical.level": 850} + ) + data = ekd.create_fieldlist([a, b]) f = AddFields(a="a", b="b", return_inputs="all") result = f.forward(data) assert len(result) == 3 for i in range(3): assert isinstance(result[i], ekd.Field) - assert {result[i].metadata("param") for i in range(2)} == {"a", "b"} - assert result[2].metadata("param") == "c" + assert {result[i].parameter.variable() for i in range(2)} == {"a", "b"} + assert result[2].parameter.variable() == "c" f = AddFields(a="a", b="b", return_inputs=("a",)) result = f.forward(data) assert len(result) == 2 for i in range(2): assert isinstance(result[i], ekd.Field) - assert result[0].metadata("param") == "a" - assert result[1].metadata("param") == "c" + assert result[0].parameter.variable() == "a" + assert result[1].parameter.variable() == "c" def test_missing_component_raises(): - a = mock_field(param="a", step=0, level=850) + a = mock_field( + **{"parameter.variable": "a", "time.valid_datetime": "2000-01-01T00:00Z", "time.step": 0, "vertical.level": 850} + ) # Missing 'b' - data = ekd.SimpleFieldList([a]) + data = ekd.create_fieldlist([a]) f = AddFields(a="a", b="b") with pytest.raises(ValueError): @@ -175,9 +182,13 @@ def backward_transform(self, a): # Missing 'b' def test_metadata_mismatch_warning(caplog): - c = mock_field(param="c", step=0, level=850) - d = mock_field(param="d", step=0, level=850) - data = ekd.SimpleFieldList([c, d]) + c = mock_field( + **{"parameter.variable": "c", "time.valid_datetime": "2000-01-01T00:00Z", "time.step": 0, "vertical.level": 850} + ) + d = mock_field( + **{"parameter.variable": "d", "time.valid_datetime": "2000-01-01T00:00Z", "time.step": 0, "vertical.level": 850} + ) + data = ekd.create_fieldlist([c, d]) f = AddFields(a="a", b="b") diff --git a/tests/test_wrappers.py b/tests/test_wrappers.py new file mode 100644 index 00000000..fcd21d11 --- /dev/null +++ b/tests/test_wrappers.py @@ -0,0 +1,164 @@ +import datetime + +import earthkit.data as ekd +import numpy as np +import pytest + +from anemoi.transform.fields import new_empty_fieldlist +from anemoi.transform.fields import new_field_from_latitudes_longitudes +from anemoi.transform.fields import new_field_from_numpy +from anemoi.transform.fields import new_field_with_metadata +from anemoi.transform.fields import new_field_with_valid_datetime +from anemoi.transform.fields import new_fieldlist_from_list + + +@pytest.fixture +def fieldlist(): + return ekd.from_source("sample", "test.grib").to_fieldlist() + + +@pytest.fixture() +def field(fieldlist): + return fieldlist[0] + + +@pytest.fixture() +def field_step_6(): + return ekd.from_source("sample", "pl.grib").to_fieldlist().sel(**{"time.step": datetime.timedelta(hours=6)})[0] + + +def test_new_fieldlist_from_list(fieldlist): + fields = list(fieldlist) + result = new_fieldlist_from_list(fields) + assert isinstance(result, ekd.FieldList) + assert len(result) == len(fields) + # ensure using the same objects (not copies) + assert all(id(f) == id(r) for f, r in zip(fields, result)) + + +def test_new_empty_fieldlist(): + result = new_empty_fieldlist() + assert isinstance(result, ekd.FieldList) + assert len(result) == 0 + + +def test_new_field_from_numpy_data_only(field): + array = field.to_numpy() + 1 + + result = new_field_from_numpy(array, template=field) + assert isinstance(result, ekd.Field) + + assert result.shape == field.shape + assert np.array_equal(result.to_numpy(), array) + + +def test_new_field_from_numpy_update_param(field): + array = field.to_numpy() + 1 + + result = new_field_from_numpy(array, template=field, param="foo") + assert isinstance(result, ekd.Field) + + assert result.parameter.variable() != field.parameter.variable() + assert result.parameter.variable() == "foo" + + assert result.shape == field.shape + assert np.array_equal(result.to_numpy(), array) + + +def test_new_field_from_numpy_update_number(field): + array = field.to_numpy() + 1 + + result = new_field_from_numpy(array, template=field, number=99) + assert isinstance(result, ekd.Field) + + assert field.ensemble.member() != result.ensemble.member() + # ensemble.member() returns a str + assert result.ensemble.member() == "99" + + assert result.shape == field.shape + assert np.array_equal(result.to_numpy(), array) + + +def test_new_field_from_numpy_update_levelist(field): + array = field.to_numpy() + 1 + + result = new_field_from_numpy(array, template=field, levelist=99) + assert isinstance(result, ekd.Field) + + assert field.vertical.level() == 0 + assert result.vertical.level() == 99 + + assert result.shape == field.shape + assert np.array_equal(result.to_numpy(), array) + + +def test_new_field_with_valid_datetime(field_step_6): + field = field_step_6 + assert field.time.step() == datetime.timedelta(hours=6) + + new_valid_datetime = field.time.valid_datetime() - field.time.step() + assert new_valid_datetime == field.time.base_datetime() + + result = new_field_with_valid_datetime(field, new_valid_datetime) + assert isinstance(result, ekd.Field) + + # check valid_datetime and step are updated (step set to 0) - base datetime unchanged + # ie. valid_datetime is the same as base_datetime + assert result.time.valid_datetime() == new_valid_datetime + assert result.time.base_datetime() == field.time.base_datetime() + assert result.time.step() != field.time.step() + assert result.time.step() == datetime.timedelta(hours=0) + + # check data unchanged + assert result.shape == field.shape + assert np.array_equal(result.to_numpy(), field.to_numpy()) + + +def test_new_field_with_metadata_update_param(field): + # new_field_with_metadata works similar to new_field_from_numpy except + # it does not allow for updating the data + result = new_field_with_metadata(field, param="foo") + assert isinstance(result, ekd.Field) + + # check param updated + assert result.parameter.variable() == "foo" + assert result.parameter.variable() != field.parameter.variable() + + # check data unchanged + assert result.shape == field.shape + assert np.array_equal(result.to_numpy(), field.to_numpy()) + + +def test_new_field_with_metadata_update_param_and_levelist(field): + # new_field_with_metadata works similar to new_field_from_numpy except + # it does not allow for updating the data + result = new_field_with_metadata(field, param="foo", levelist=99) + assert isinstance(result, ekd.Field) + + # check param and level updated + assert result.parameter.variable() == "foo" + assert result.parameter.variable() != field.parameter.variable() + assert field.vertical.level() == 0 + assert result.vertical.level() == 99 + + # check data unchanged + assert result.shape == field.shape + assert np.array_equal(result.to_numpy(), field.to_numpy()) + + +def test_new_field_from_latitudes_longitudes(field): + lat, lon = field.geography.latlons() + new_lat = lat + 12 + new_lon = lon - 34 + + result = new_field_from_latitudes_longitudes(field, new_lat, new_lon) + assert isinstance(result, ekd.Field) + + # check grid points updated + result_lat, result_lon = result.geography.latlons() + assert np.array_equal(result_lat, new_lat) + assert np.array_equal(result_lon, new_lon) + + # check data unchanged + assert result.shape == field.shape + assert np.array_equal(result.to_numpy(), field.to_numpy()) diff --git a/tests/utils/__init__.py b/tests/utils/__init__.py index aeb3467f..b046c70f 100644 --- a/tests/utils/__init__.py +++ b/tests/utils/__init__.py @@ -7,6 +7,8 @@ # granted to it by virtue of its status as an intergovernmental organisation # nor does it submit to any jurisdiction. from collections import defaultdict +from collections.abc import Mapping +from typing import Any import earthkit.data as ekd import numpy as np @@ -18,32 +20,39 @@ def collect_fields_by_param(pipeline): fields = defaultdict(list) for field in pipeline: - param = field.metadata("param") + param = field.parameter.variable() fields[param].append(field) return fields def assert_fields_equal(field_a, field_b, exclude_keys=None): - METADATA_KEYS = ["param", "valid_datetime", "latitudes", "longitudes", "levelist"] + METADATA_KEYS = [ + "parameter.variable", + "time.valid_datetime", + "geography.distinct_latitudes", + "geography.distinct_longitudes", + "vertical.level", + ] if exclude_keys is None: exclude_keys = [] exclude_keys = set(exclude_keys) + # TODO: remove this? # workaround for unreliable __contains__ in potentially wrapped objects def metadata_contains(field, key): try: - field.metadata(key) + field.get(key) return True except KeyError: return False for key in set(METADATA_KEYS) - exclude_keys: try: - assert field_a.metadata(key) == field_b.metadata(key) + assert field_a.get(key) == field_b.get(key) except ValueError: # if ValueError, assume not just scalar values - use numpy for comparison - assert np.allclose(field_a.metadata(key), field_b.metadata(key)) + assert np.allclose(field_a.get(key), field_b.get(key)) except KeyError: in_a = metadata_contains(field_a, key) in_b = metadata_contains(field_b, key) @@ -64,7 +73,7 @@ def __init__(self, fields, params=None): def forward(self, *args, **kwargs): fields = [] for f in self._fields: - if self.params and f.metadata("param") in self.params: + if self.params and f.parameter.variable() in self.params: fields.append(f) return new_fieldlist_from_list(fields) @@ -83,11 +92,11 @@ def forward(self, *args, **kwargs): fields = [] params = [] for f in self._fields: - if self.params and f.metadata("param") in self.params: + if self.params and f.parameter.variable() in self.params: fields.append(f) - params.append(f.metadata("param")) + params.append(f.parameter.variable()) for f in self._additional_fields: - if self.params and f.metadata("param") in self.params and f.metadata("param") not in params: + if self.params and f.parameter.variable() in self.params and f.parameter.variable() not in params: fields.append(f) return new_fieldlist_from_list(fields) @@ -105,10 +114,49 @@ def compare_npz_files(file1, file2): def mock_field(**metadata): - class MetadataOverride(ekd.core.metadata.RawMetadata): - def as_namespace(self, namespace): - if namespace != "mars": - raise ValueError(f"Unknown namespace {namespace}") - return dict(self) + field_spec = {"data.values": np.array([1])} | metadata + field_spec = group_component_dict(field_spec) + return ekd.from_source("list-of-dicts", [field_spec]).to_fieldlist()[0] - return ekd.ArrayField(array=[1], metadata=MetadataOverride(**metadata)) + +def group_component_dict(components: Mapping[str, Any]) -> dict[str, dict[str, Any]]: + """Groups dictionaries in the form {'x.y': 'u', 'x.z': 'v', ...} into {'x': {'y': 'v', 'z': 'v'}, ...}""" + + SEP = "." + result = {} + + for key, value in components.items(): + if not isinstance(key, str): + raise TypeError(f"Expected key to be a str, got {type(key)}: {key!r}") + + if key.startswith(SEP): + raise ValueError(f"Invalid key {key}: cannot start with '{SEP}'") + + head, found, tail = key.partition(SEP) + + if not found: + # key does not have components - must be a full dict + if not isinstance(value, dict): + raise ValueError(f"Value of key {key} must be a dict, got {type(value)}") + + if head in result: + raise ValueError(f"Duplicate key: {key}") + result[head] = value + continue + + # sep was found - therefore key is like "x.", i.e. with no tail + if not tail: + raise ValueError(f"Invalid key {key}: empty tail after '{SEP}'") + + # key is in the form "x.y" (assume two levels max) + if SEP in tail: + raise ValueError(f"Invalid key: {key}, cannot have more than one '{SEP}'") + + if head not in result: + result[head] = {} + + if tail in result[head]: + raise KeyError(f"Duplicate key: {key} already exists") + + result[head][tail] = value + return result