diff --git a/.github/ci/recipe.yaml b/.github/ci/recipe.yaml index 36212e131..fe9f08dfa 100644 --- a/.github/ci/recipe.yaml +++ b/.github/ci/recipe.yaml @@ -46,6 +46,7 @@ requirements: - holoviews >= 1.22.0 # https://github.com/prefix-dev/rattler-build/issues/2326 - pooch >=1.8.0 - polars >=1.31.0 + - tabulate >=0.10.0 tests: - python: diff --git a/pixi.toml b/pixi.toml index 4b1c05ea3..6ff79db06 100644 --- a/pixi.toml +++ b/pixi.toml @@ -35,6 +35,8 @@ cf_xarray = ">=0.8.6" cftime = ">=1.6.3" pooch = ">=1.8.0" polars = ">=1.31.0" +tabulate = ">=0.10.0" + [dependencies] parcels = { path = "." } @@ -155,6 +157,7 @@ numpydoc-lint = { cmd = "python tools/numpydoc-public-api.py", description = "Li mypy = "*" lxml = "*" # in CI types-tqdm = "*" +pandas-stubs = "*" [feature.typing.tasks] typing = { cmd = "mypy src/parcels --install-types", description = "Run static type checking with mypy." } @@ -177,3 +180,6 @@ test-notebooks = { features = ["test", "notebooks"], solve-group = "main" } docs = { features = ["docs", "notebooks"], solve-group = "docs" } typing = { features = ["typing"], solve-group = "main" } build = { features = ["rattler-build"] } + +[pypi-dependencies] +detect-test-pollution = ">=1.2.0, <2" diff --git a/src/parcels/_core/basegrid.py b/src/parcels/_core/basegrid.py index 57da1120f..ce4f88e3f 100644 --- a/src/parcels/_core/basegrid.py +++ b/src/parcels/_core/basegrid.py @@ -8,6 +8,7 @@ import numpy as np +import parcels._typing as ptyping from parcels._core.spatialhash import SpatialHash if TYPE_CHECKING: @@ -25,6 +26,7 @@ class BaseGrid(ABC): """Base class for parcels.XGrid and parcels.UxGrid defining common methods and properties""" _spatialhash: SpatialHash | None + _mesh: ptyping.Mesh @abstractmethod def search(self, z: float, y: float, x: float, ei=None) -> dict[str, tuple[int, float | np.ndarray]]: diff --git a/src/parcels/_core/fieldset.py b/src/parcels/_core/fieldset.py index 45310f2cb..bf7f9df59 100644 --- a/src/parcels/_core/fieldset.py +++ b/src/parcels/_core/fieldset.py @@ -21,6 +21,7 @@ from parcels._core.utils.time import get_datetime_type_calendar from parcels._core.utils.time import is_compatible as datetime_is_compatible from parcels._python import NOTSET, NotSetType +from parcels._reprs import fieldset_describe from parcels.interpolators import ( XConstantField, ) @@ -284,6 +285,10 @@ def from_sgrid_conventions( model = StructuredModelData.from_sgrid_conventions(ds, mesh, vector_fields) return cls([model]) + def describe(self): + """Return a table description of a FieldSet, which fields it has and their interpolation methods.""" + return fieldset_describe(self) + def assert_compatible_fieldsets(left: FieldSet, right: FieldSet) -> None: """Assert that two FieldSets can be combined without name conflicts. diff --git a/src/parcels/_reprs.py b/src/parcels/_reprs.py index bc23ad0cd..a5e60d9da 100644 --- a/src/parcels/_reprs.py +++ b/src/parcels/_reprs.py @@ -3,14 +3,19 @@ from __future__ import annotations import textwrap -from typing import TYPE_CHECKING, Any, cast +from dataclasses import dataclass +from typing import TYPE_CHECKING, Any, Literal, cast import numpy as np import xarray as xr +from parcels._python import isinstance_noimport + if TYPE_CHECKING: from parcels import Field, FieldSet, ParticleSet from parcels._core.field import VectorField + from parcels._core.model import ModelData + from parcels._core.utils.time import TimeInterval def fieldset_repr(fieldset: FieldSet) -> str: @@ -177,3 +182,81 @@ def _format_list_items_multiline(items: list[str] | dict, level: int = 1, with_b def is_builtin_object(obj): return obj.__class__.__module__ == "builtins" + + +@dataclass +class _FieldSetDescriptionRow: + type_: Literal["Field", "VectorField", "Context"] + model_id: int | None + name: str + interp_method_or_value: str + + def to_dict(self) -> dict[str, str]: + return { + "Name": self.name, + "Type": self.type_, + "Grid number": str(self.model_id) if self.model_id is not None else "-", + "Interp method / value": self.interp_method_or_value, + } + + +def _print_table(rows: list[_FieldSetDescriptionRow]) -> str: + import pandas as pd + + dicts = [r.to_dict() for r in rows] + return pd.DataFrame(dicts).sort_values(["Grid number", "Type", "Name"]).to_markdown(index=False) + + +def _print_time_interval(time_interval: TimeInterval | None) -> str: + if time_interval is None: + return repr(time_interval) + return repr((time_interval.left, time_interval.right)) + + +def fieldset_describe(fieldset: FieldSet) -> str: + rows: list[_FieldSetDescriptionRow] = [] + models: dict[int, int] = {} # mapping of memory ID to a human readable ID + + assert fieldset._fields is not None + + for field in fieldset._fields.values(): + model_id: int + + # Set human readable model ID + parent_id = id(_get_parent_model(field)) + models[parent_id] = models.get(parent_id, len(models)) + model_id = models[parent_id] + + type_ = cast(Literal["Field", "VectorField", "Context"], field.__class__.__name__) + + rows.append( + _FieldSetDescriptionRow( + type_=type_, + model_id=model_id, + name=field.name, + interp_method_or_value=repr(field.interp_method), + ) + ) + for k, v in fieldset.context.items(): + rows.append( + _FieldSetDescriptionRow( + type_="Context", + model_id=None, + name=k, + interp_method_or_value=repr(v), + ) + ) + return ( + _print_table(rows) + + f"""\ + + +mesh: {fieldset.models[0].grid._mesh} +time interval: {_print_time_interval(fieldset.time_interval)}""" + ) + + +def _get_parent_model(field: Field | VectorField) -> ModelData: + if isinstance_noimport(field, "Field"): + return field.model # type:ignore[union-attr] + return field.U.model # type:ignore[union-attr] diff --git a/src/parcels/interpolators/_base.py b/src/parcels/interpolators/_base.py index a4dbaf5e0..6a550b240 100644 --- a/src/parcels/interpolators/_base.py +++ b/src/parcels/interpolators/_base.py @@ -7,8 +7,14 @@ class ScalarInterpolator(ABC): def interp(self, particle_positions, grid_positions, field) -> Any: #! API a WIP ... + def __repr__(self): + return f"{self.__class__.__name__}(...)" + class VectorInterpolator(ABC): @abstractmethod def interp(self, particle_positions, grid_positions, vectorfield) -> Any: #! API a WIP ... + + def __repr__(self): + return f"{self.__class__.__name__}(...)" diff --git a/tests/test_fieldset.py b/tests/test_fieldset.py index 799517a06..8cf7d4661 100644 --- a/tests/test_fieldset.py +++ b/tests/test_fieldset.py @@ -18,6 +18,21 @@ ds = datasets_structured["ds_2d_left"] +@pytest.fixture +def fieldset_two_models(): + ds1 = datasets_structured["ds_2d_left"][["U_A_grid", "V_A_grid", "grid"]].rename({"U_A_grid": "U", "V_A_grid": "V"}) + ds2 = datasets_structured["ds_2d_left"][["U_A_grid", "V_A_grid", "grid"]].rename( + {"U_A_grid": "U_wind", "V_A_grid": "V_wind"} + ) + + fset1 = FieldSet.from_sgrid_conventions(ds1, mesh="flat") + fset2 = FieldSet.from_sgrid_conventions(ds2, mesh="flat", vector_fields={"UV_wind": ("U_wind", "V_wind")}) + fset2.add_context("my_value", 2.0) + fset2.add_context("my_list", [1, 2, "hello"]) + fset2.add_constant_field("constant_field", 3.0) + return fset1 + fset2 + + def test_fieldset_init_wrong_types(): with pytest.raises(ValueError, match="Expected `model` to be a ModelData object. Got .*"): FieldSet([1.0, 2.0, 3.0]) @@ -373,3 +388,27 @@ def test_fieldset_add_context_values(): assert fset.context["c1"] == 1.0 assert fset.context["c2"] == 2.0 + + +@pytest.mark.xfail( + reason="There's test pollution occuring between test_fieldKh_Brownian and this test due to how constant fields are handled. We should remove this global state." +) +def test_fieldset_describe(fieldset_two_models: FieldSet): + fieldset = fieldset_two_models + expected = """\ +| Name | Type | Grid number | Interp method / value | +|:---------------|:------------|:--------------|:------------------------| +| my_list | Context | - | [1, 2, 'hello'] | +| my_value | Context | - | 2.0 | +| U | Field | 0 | XLinear(...) | +| V | Field | 0 | XLinear(...) | +| UV | VectorField | 0 | XLinear_Velocity(...) | +| U_wind | Field | 1 | XLinear(...) | +| V_wind | Field | 1 | XLinear(...) | +| UV_wind | VectorField | 1 | XLinear_Velocity(...) | +| constant_field | Field | 2 | XConstantField(...) | + +mesh: flat +time interval: (np.datetime64('2000-01-01T00:00:00.000000000'), np.datetime64('2001-01-01T00:00:00.000000000'))""" + actual = fieldset.describe() + assert actual == expected