diff --git a/src/parcels/__init__.py b/src/parcels/__init__.py index 78227247b..81c901ca9 100644 --- a/src/parcels/__init__.py +++ b/src/parcels/__init__.py @@ -22,6 +22,7 @@ from parcels._core.basegrid import BaseGrid from parcels._core.uxgrid import UxGrid from parcels._core.xgrid import XGrid +from parcels._core.mesh import SphericalMesh from parcels._core.statuscodes import ( AllParcelsErrorCodes, @@ -54,6 +55,7 @@ "BaseGrid", "UxGrid", "XGrid", + "SphericalMesh", # Status codes and errors "AllParcelsErrorCodes", "FieldInterpolationError", diff --git a/src/parcels/_core/basegrid.py b/src/parcels/_core/basegrid.py index ce4f88e3f..7b7c7e9c4 100644 --- a/src/parcels/_core/basegrid.py +++ b/src/parcels/_core/basegrid.py @@ -8,7 +8,7 @@ import numpy as np -import parcels._typing as ptyping +from parcels._core.mesh import FlatMesh, SphericalMesh from parcels._core.spatialhash import SpatialHash if TYPE_CHECKING: @@ -26,7 +26,7 @@ class BaseGrid(ABC): """Base class for parcels.XGrid and parcels.UxGrid defining common methods and properties""" _spatialhash: SpatialHash | None - _mesh: ptyping.Mesh + _mesh: FlatMesh | SphericalMesh @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 232ca2724..5864b9288 100644 --- a/src/parcels/_core/fieldset.py +++ b/src/parcels/_core/fieldset.py @@ -189,7 +189,7 @@ def to_windowed_arrays(self, *, max_levels: int | None = None): model.to_windowed_arrays(max_levels=max_levels) return self - def add_constant_field(self, name: str, value, mesh: ptyping.Mesh = "spherical"): + def add_constant_field(self, name: str, value, mesh: ptyping.TMesh = "spherical"): """Wrapper function to add a Field that is constant in space, useful e.g. when using constant horizontal diffusivity @@ -287,7 +287,7 @@ def from_ugrid_conventions( def from_sgrid_conventions( cls, ds: xr.Dataset, - mesh: ptyping.Mesh | None = None, + mesh: ptyping.TMesh | None = None, vector_fields: ptyping.VectorFields | NotSetType = NOTSET, ): # TODO: Update mesh to be discovered from the dataset metadata """Create a FieldSet from a dataset using SGRID convention metadata. diff --git a/src/parcels/_core/index_search.py b/src/parcels/_core/index_search.py index ebd4f668a..77e43aff3 100644 --- a/src/parcels/_core/index_search.py +++ b/src/parcels/_core/index_search.py @@ -109,7 +109,7 @@ def curvilinear_point_in_cell(grid, y: np.ndarray, x: np.ndarray, yi: np.ndarray dtype=float, ) - if grid._mesh == "spherical": + if grid._mesh.is_spherical(): xsi, eta = _bilinear_inverse_tangent_plane(clon, clat, x, y) is_in_cell = np.where((xsi >= 0) & (xsi <= 1) & (eta >= 0) & (eta <= 1), 1, 0) else: @@ -318,7 +318,7 @@ def uxgrid_point_in_cell(grid, y: np.ndarray, x: np.ndarray, yi: np.ndarray, xi: coords : np.ndarray Barycentric coordinates of the points within their respective cells. """ - if grid._mesh == "spherical": + if grid._mesh.is_spherical(): lon_rad = np.deg2rad(x) lat_rad = np.deg2rad(y) x_cart, y_cart, z_cart = _latlon_rad_to_xyz(lat_rad, lon_rad) diff --git a/src/parcels/_core/kernel.py b/src/parcels/_core/kernel.py index f50c2876e..f68c8bf09 100644 --- a/src/parcels/_core/kernel.py +++ b/src/parcels/_core/kernel.py @@ -141,9 +141,9 @@ def check_fieldsets_in_kernels(self, kernel): # TODO v4: this can go into anoth stacklevel=2, ) self.fieldset.add_context("RK45_tol", 10) - if self.fieldset.U.grid._mesh == "spherical": + if self.fieldset.U.grid._mesh.is_spherical(): self.fieldset.RK45_tol /= ( - 1852 * 60 + self.fieldset.U.grid.deg2m ) # TODO does not account for zonal variation in meter -> degree conversion if not hasattr(self.fieldset, "RK45_min_dt"): warnings.warn( diff --git a/src/parcels/_core/mesh.py b/src/parcels/_core/mesh.py new file mode 100644 index 000000000..b60b0d6c8 --- /dev/null +++ b/src/parcels/_core/mesh.py @@ -0,0 +1,73 @@ +from abc import ABC, abstractmethod +from typing import Literal + +import numpy as np + +EARTH_RADIUS = 6366707.019493707 + + +class BaseMesh(ABC): + radius: float | None + + @abstractmethod + def is_spherical(self) -> bool: ... + + +class SphericalMesh(BaseMesh): + """Spherical mesh object with configurable planetary radius. + + Pass to FieldSet object as ``mesh=SphericalMesh(radius=...)``. + radius is in meters; defaults to Earth radius. + """ + + def __init__(self, radius: float = EARTH_RADIUS): + if not isinstance(radius, (int, float, np.number)): + raise TypeError(f"radius must be a number, got {type(radius).__name__}") + if radius <= 0: + raise ValueError(f"radius must be positive, got {radius}") + self.radius = radius + + @property + def deg2m(self) -> float: + """Meters per degree of arc.""" + assert self.radius is not None + return self.radius * np.pi / 180.0 + + def is_spherical(self): + return True + + def __repr__(self) -> str: + return f"SphericalMesh(radius={self.radius})" + + +class FlatMesh(BaseMesh): + """Flat mesh object.""" + + def __init__(self): + self.radius = None + return + + def __repr__(self) -> str: + return "FlatMesh()" + + def is_spherical(self): + return False + + +TMesh = SphericalMesh | Literal["spherical", "flat"] # corresponds with `mesh` + + +def get_mesh(mesh: TMesh): + if isinstance(mesh, SphericalMesh): + return mesh + if mesh == "flat": + return FlatMesh() + if mesh == "spherical": + return SphericalMesh(EARTH_RADIUS) + raise ValueError(f"mesh must be 'flat', 'spherical', or a SphericalMesh object. Got {mesh=!r}") + + +def is_spherical(mesh: FlatMesh | SphericalMesh): + if isinstance(mesh, SphericalMesh): + return True + return False diff --git a/src/parcels/_core/model.py b/src/parcels/_core/model.py index 79c9c6460..24542f7ba 100644 --- a/src/parcels/_core/model.py +++ b/src/parcels/_core/model.py @@ -23,7 +23,6 @@ ) from parcels._logger import logger from parcels._python import NOTSET, NotSetType -from parcels._typing import Mesh from parcels.convert import _ds_rename_using_standard_names from parcels.interpolators import ( CGrid_Velocity, @@ -133,7 +132,7 @@ def preprocess_sgrid_model_data(ds: xr.Dataset) -> xr.Dataset: class StructuredModelData(ModelData): - def __init__(self, data: xr.Dataset, mesh: Mesh, vector_field_components: ptyping.VectorFields): + def __init__(self, data: xr.Dataset, mesh: ptyping.TMesh, vector_field_components: ptyping.VectorFields): if not isinstance(data, xr.Dataset): raise ValueError(f"Expected `data` to be an xarray.Dataset . Got {type(data)}") @@ -182,7 +181,7 @@ def construct_fields(self) -> list[Field | VectorField]: @classmethod def from_sgrid_conventions( - cls, ds: xr.Dataset, mesh: Mesh | None, vector_fields: ptyping.VectorFields | NotSetType + cls, ds: xr.Dataset, mesh: ptyping.TMesh | None, vector_fields: ptyping.VectorFields | NotSetType ) -> Self: ds = ds.copy() if mesh is None: @@ -329,7 +328,9 @@ def scalar_field_names(self) -> list[str]: return list(self.data.data_vars) @classmethod - def from_ugrid_conventions(cls, ds: ux.UxDataset, mesh: Mesh, vector_fields: ptyping.VectorFields | NotSetType): + def from_ugrid_conventions( + cls, ds: ux.UxDataset, mesh: ptyping.TMesh, vector_fields: ptyping.VectorFields | NotSetType + ): ds_dims = list(ds.dims) if not all(dim in ds_dims for dim in ["time", "zf", "zc"]): raise ValueError( @@ -353,7 +354,7 @@ def from_ugrid_conventions(cls, ds: ux.UxDataset, mesh: Mesh, vector_fields: pty # TODO: Refactor later into something like `parcels._metadata.discover(dataset)` helper that can be used to discover important metadata like this. I think this whole metadata handling should be refactored into its own module. -def _get_mesh_type_from_sgrid_dataset(ds_sgrid: xr.Dataset) -> Mesh: +def _get_mesh_type_from_sgrid_dataset(ds_sgrid: xr.Dataset) -> ptyping.TMesh: """Small helper to inspect SGRID metadata and dataset metadata to determine mesh type.""" sgrid_metadata = ds_sgrid.sgrid.metadata diff --git a/src/parcels/_core/particlefile.py b/src/parcels/_core/particlefile.py index 64eb45e66..50bdf309c 100644 --- a/src/parcels/_core/particlefile.py +++ b/src/parcels/_core/particlefile.py @@ -14,6 +14,7 @@ import xarray as xr import parcels +from parcels._core.mesh import BaseMesh from parcels._core.particle import ParticleClass from parcels._core.particlesetview import ParticleSetView from parcels._core.utils.time import timedelta_to_float @@ -119,14 +120,14 @@ def __init__( def __repr__(self) -> str: return particlefile_repr(self) - def set_metadata(self, parcels_grid_mesh: Literal["spherical", "flat"]): + def set_metadata(self, parcels_grid_mesh: BaseMesh): self.metadata.update( { "feature_type": "trajectory", "Conventions": "CF-1.6/CF-1.7", "ncei_template_version": "NCEI_NetCDF_Trajectory_Template_v2.0", "parcels_version": parcels.__version__, - "parcels_grid_mesh": parcels_grid_mesh, + "parcels_grid_mesh": repr(parcels_grid_mesh), } ) diff --git a/src/parcels/_core/spatialhash.py b/src/parcels/_core/spatialhash.py index 7d0696859..bcb11914a 100644 --- a/src/parcels/_core/spatialhash.py +++ b/src/parcels/_core/spatialhash.py @@ -55,7 +55,7 @@ def __init__( if isinstance_noimport(grid, "XGrid"): self._coord_dim = 2 # Number of computational coordinates is 2 (bilinear interpolation) - if self._source_grid._mesh == "spherical": + if self._source_grid._mesh.is_spherical(): lon = np.deg2rad(self._source_grid.lon) lat = np.deg2rad(self._source_grid.lat) x, y, z = _latlon_rad_to_xyz(lat, lon) @@ -160,7 +160,7 @@ def __init__( elif isinstance_noimport(grid, "UxGrid"): self._coord_dim = grid.uxgrid.n_max_face_nodes # Number of barycentric coordinates - if self._source_grid._mesh == "spherical": + if self._source_grid._mesh.is_spherical(): # Reshape node coordinates to (nfaces, nnodes_per_face) nids = self._source_grid.uxgrid.face_node_connectivity.values lon = self._source_grid.uxgrid.node_lon.values[nids] @@ -404,7 +404,7 @@ def query(self, y, x): y = np.asarray(y) x = np.asarray(x) - if self._source_grid._mesh == "spherical": + if self._source_grid._mesh.is_spherical(): # Convert coords to Cartesian coordinates (x, y, z) lat = np.deg2rad(y) lon = np.deg2rad(x) diff --git a/src/parcels/_core/utils/interpolation.py b/src/parcels/_core/utils/interpolation.py index a48c0ff21..5900c2380 100644 --- a/src/parcels/_core/utils/interpolation.py +++ b/src/parcels/_core/utils/interpolation.py @@ -3,7 +3,7 @@ import numpy as np -from parcels._typing import Mesh +from parcels._core.mesh import BaseMesh __all__ = [] @@ -61,12 +61,12 @@ def dphidxsi3D_lin(zeta: float, eta: float, xsi: float) -> tuple[list[float], li def dxdxsi3D_lin( - hexa_z: list[float], hexa_y: list[float], hexa_x: list[float], zeta: float, eta: float, xsi: float, mesh: Mesh + hexa_z: list[float], hexa_y: list[float], hexa_x: list[float], zeta: float, eta: float, xsi: float, mesh: BaseMesh, + deg2m: float = 1852 * 60.0 ) -> tuple[float, float, float, float, float, float, float, float, float]: dphidxsi, dphideta, dphidzet = dphidxsi3D_lin(zeta, eta, xsi) - if mesh == 'spherical': - deg2m = 1852 * 60. + if mesh.is_spherical(): rad = np.pi / 180. lat = (1-xsi) * (1-eta) * hexa_y[0] + \ xsi * (1-eta) * hexa_y[1] + \ @@ -92,9 +92,10 @@ def dxdxsi3D_lin( def jacobian3D_lin( - hexa_z: list[float], hexa_y: list[float], hexa_x: list[float], zeta: float, eta: float, xsi: float, mesh: Mesh + hexa_z: list[float], hexa_y: list[float], hexa_x: list[float], zeta: float, eta: float, xsi: float, mesh: BaseMesh, + deg2m: float = 1852 * 60.0 ) -> float: - dxdxsi, dxdeta, dxdzet, dydxsi, dydeta, dydzet, dzdxsi, dzdeta, dzdzet = dxdxsi3D_lin(hexa_z, hexa_y, hexa_x, zeta, eta, xsi, mesh) + dxdxsi, dxdeta, dxdzet, dydxsi, dydeta, dydzet, dzdxsi, dzdeta, dzdzet = dxdxsi3D_lin(hexa_z, hexa_y, hexa_x, zeta, eta, xsi, mesh, deg2m) jac = ( dxdxsi * (dydeta * dzdzet - dzdeta * dydzet) @@ -112,7 +113,7 @@ def jacobian3D_lin_face( eta: float, xsi: float, orientation: Literal["zonal", "meridional", "vertical"], - mesh: Mesh, + mesh: BaseMesh, ) -> float: dxdxsi, dxdeta, dxdzet, dydxsi, dydeta, dydzet, dzdxsi, dzdeta, dzdzet = dxdxsi3D_lin(hexa_z, hexa_y, hexa_x, zeta, eta, xsi, mesh) @@ -174,10 +175,11 @@ def interpolate(phi: Callable[[float], list[float]], f: list[float], xsi: float) return np.dot(phi(xsi), f) -def _geodetic_distance(lat1: float, lat2: float, lon1: float, lon2: float, mesh: Mesh, lat: float) -> float: - if mesh == "spherical": +def _geodetic_distance( + lat1: float, lat2: float, lon1: float, lon2: float, mesh: BaseMesh, lat: float, deg2m: float = 1852 * 60.0 +) -> float: + if mesh.is_spherical(): rad = np.pi / 180.0 - deg2m = 1852 * 60.0 return np.sqrt(((lon2 - lon1) * deg2m * np.cos(rad * lat)) ** 2 + ((lat2 - lat1) * deg2m) ** 2) else: return np.sqrt((lon2 - lon1) ** 2 + (lat2 - lat1) ** 2) diff --git a/src/parcels/_core/uxgrid.py b/src/parcels/_core/uxgrid.py index 39201c2e8..7f9360ab9 100644 --- a/src/parcels/_core/uxgrid.py +++ b/src/parcels/_core/uxgrid.py @@ -7,7 +7,7 @@ from parcels._core.basegrid import BaseGrid from parcels._core.index_search import GRID_SEARCH_ERROR, _search_1d_array, uxgrid_point_in_cell -from parcels._typing import assert_valid_mesh +from parcels._core.mesh import SphericalMesh, get_mesh _UXGRID_AXES = Literal["Z", "FACE"] @@ -18,7 +18,9 @@ class UxGrid(BaseGrid): for interpolation on unstructured grids. """ - def __init__(self, grid: ux.grid.Grid, z: ux.UxDataArray, mesh) -> None: + def __init__( + self, grid: ux.grid.Grid, z: ux.UxDataArray, mesh: Literal["flat", "spherical"] | SphericalMesh + ) -> None: """ Initializes the UxGrid with a uxarray grid and vertical coordinate array. @@ -41,11 +43,9 @@ def __init__(self, grid: ux.grid.Grid, z: ux.UxDataArray, mesh) -> None: if z.ndim != 1: raise ValueError("z must be a 1D array of vertical coordinates") self.z = z - self._mesh = mesh + self._mesh = get_mesh(mesh) self._spatialhash = None - assert_valid_mesh(mesh) - @property def depth(self): """ @@ -73,6 +73,13 @@ def get_axis_dim(self, axis: _UXGRID_AXES) -> int: elif axis == "FACE": return self.uxgrid.n_face + @property + def deg2m(self) -> float: + """Metres per arcdegree for this grid's mesh.""" + if not self._mesh.is_spherical(): + return 1.0 + return self._mesh.deg2m + def search(self, z, y, x, ei=None, tol=1e-6): """ Search for the grid cell (face) and vertical layer that contains the given points. diff --git a/src/parcels/_core/xgrid.py b/src/parcels/_core/xgrid.py index 36a908475..94724e207 100644 --- a/src/parcels/_core/xgrid.py +++ b/src/parcels/_core/xgrid.py @@ -13,6 +13,7 @@ import parcels._typing as ptyping from parcels._core.basegrid import BaseGrid from parcels._core.index_search import _search_1d_array, _search_indices_curvilinear_2d +from parcels._core.mesh import SphericalMesh, get_mesh from parcels._sgrid.accessor import _get_dim_to_axis_mapping from parcels._sgrid.core import SGRID_PADDING_TO_XGCM_POSITION @@ -165,12 +166,12 @@ class XGrid(BaseGrid): """ - def __init__(self, model_data: xr.Dataset, mesh): + def __init__(self, model_data: xr.Dataset, mesh: Literal["flat", "spherical"] | SphericalMesh): self.sgrid_metadata = model_data.sgrid.metadata self._ds = model_data grid = XgcmLikeGrid(self.sgrid_metadata, model_data) self.xgcm_grid = grid - self._mesh = mesh + self._mesh = get_mesh(mesh) self._spatialhash = None ds = model_data @@ -186,7 +187,6 @@ def __init__(self, model_data: xr.Dataset, mesh): if "Z" in grid.axes: assert_valid_depth(ds["depth"]) - ptyping.assert_valid_mesh(mesh) self._ds = ds # def __repr__(self): @@ -256,6 +256,13 @@ def _datetimes(self): def time(self): return self._datetimes.astype(np.float64) / 1e9 + @property + def deg2m(self) -> float: + """Metres per degree of arc for this grid's mesh.""" + if not self._mesh.is_spherical(): + return 1.0 + return self._mesh.deg2m + @cached_property def xdim(self) -> int: return self.get_axis_dim("X") diff --git a/src/parcels/_typing.py b/src/parcels/_typing.py index e8993cb05..87e23a53c 100644 --- a/src/parcels/_typing.py +++ b/src/parcels/_typing.py @@ -9,11 +9,13 @@ import os from collections.abc import Callable, Mapping from datetime import datetime -from typing import TYPE_CHECKING, Any, Literal, get_args +from typing import TYPE_CHECKING, Literal, get_args import numpy as np from cftime import datetime as cftime_datetime +from parcels._core.mesh import TMesh # noqa: F401 + if TYPE_CHECKING: import xgcm @@ -33,7 +35,6 @@ InterpMethodOption | dict[str, InterpMethodOption] ) # corresponds with `interp_method` (which can also be dict mapping field names to method) PathLike = str | os.PathLike -Mesh = Literal["spherical", "flat"] # corresponds with `mesh` VectorType = Literal["3D", "3DSigma", "2D"] | None # corresponds with `vector_type` GridIndexingType = Literal["pop", "mom5", "mitgcm", "nemo", "croco"] # corresponds with `gridindexingtype` NetcdfEngine = Literal["netcdf4", "xarray", "scipy"] @@ -69,8 +70,3 @@ def _validate_against_pure_literal(value, typing_literal): if value not in get_args(typing_literal): msg = f"Invalid value {value!r}. Valid options are {get_args(typing_literal)!r}" raise ValueError(msg) - - -# Assertion functions to clean user input -def assert_valid_mesh(value: Any): - _validate_against_pure_literal(value, Mesh) diff --git a/src/parcels/interpolators/_uxinterpolators.py b/src/parcels/interpolators/_uxinterpolators.py index 80e804475..17249bfdc 100644 --- a/src/parcels/interpolators/_uxinterpolators.py +++ b/src/parcels/interpolators/_uxinterpolators.py @@ -170,9 +170,9 @@ def interp( ): u = vectorfield.U.interp_method.interp(particle_positions, grid_positions, vectorfield.U) v = vectorfield.V.interp_method.interp(particle_positions, grid_positions, vectorfield.V) - if vectorfield.grid._mesh == "spherical": - u /= 1852 * 60 * np.cos(np.deg2rad(particle_positions["y"])) - v /= 1852 * 60 + if vectorfield.grid._mesh.is_spherical(): + u /= vectorfield.grid.deg2m * np.cos(np.deg2rad(particle_positions["y"])) + v /= vectorfield.grid.deg2m if "3D" in vectorfield.vector_type: w = vectorfield.W.interp_method.interp(particle_positions, grid_positions, vectorfield.W) diff --git a/src/parcels/interpolators/_xinterpolators.py b/src/parcels/interpolators/_xinterpolators.py index 388925d1a..d89ed9cff 100644 --- a/src/parcels/interpolators/_xinterpolators.py +++ b/src/parcels/interpolators/_xinterpolators.py @@ -148,9 +148,9 @@ def interp( _xlinear = XLinear() u = _xlinear.interp(particle_positions, grid_positions, vectorfield.U) v = _xlinear.interp(particle_positions, grid_positions, vectorfield.V) - if vectorfield.grid._mesh == "spherical": - u /= 1852 * 60 * np.cos(np.deg2rad(particle_positions["y"])) - v /= 1852 * 60 + if vectorfield.grid._mesh.is_spherical(): + u /= vectorfield.grid.deg2m * np.cos(np.deg2rad(particle_positions["y"])) + v /= vectorfield.grid.deg2m if vectorfield.W: w = _xlinear.interp(particle_positions, grid_positions, vectorfield.W) @@ -196,21 +196,21 @@ def interp( px = np.array([grid.lon[yi, xi], grid.lon[yi, xi + 1], grid.lon[yi + 1, xi + 1], grid.lon[yi + 1, xi]]) py = np.array([grid.lat[yi, xi], grid.lat[yi, xi + 1], grid.lat[yi + 1, xi + 1], grid.lat[yi + 1, xi]]) - if grid._mesh == "spherical": + if grid._mesh.is_spherical(): px = ((px + 180.0) % 360.0) - 180.0 px[1:] = np.where(px[1:] - px[0] > 180, px[1:] - 360, px[1:]) px[1:] = np.where(-px[1:] + px[0] > 180, px[1:] + 360, px[1:]) c1 = i_u._geodetic_distance( - py[0], py[1], px[0], px[1], grid._mesh, np.einsum("ij,ji->i", i_u.phi2D_lin(0.0, xsi), py) + py[0], py[1], px[0], px[1], grid._mesh, np.einsum("ij,ji->i", i_u.phi2D_lin(0.0, xsi), py), grid.deg2m ) c2 = i_u._geodetic_distance( - py[1], py[2], px[1], px[2], grid._mesh, np.einsum("ij,ji->i", i_u.phi2D_lin(eta, 1.0), py) + py[1], py[2], px[1], px[2], grid._mesh, np.einsum("ij,ji->i", i_u.phi2D_lin(eta, 1.0), py), grid.deg2m ) c3 = i_u._geodetic_distance( - py[2], py[3], px[2], px[3], grid._mesh, np.einsum("ij,ji->i", i_u.phi2D_lin(1.0, xsi), py) + py[2], py[3], px[2], px[3], grid._mesh, np.einsum("ij,ji->i", i_u.phi2D_lin(1.0, xsi), py), grid.deg2m ) c4 = i_u._geodetic_distance( - py[3], py[0], px[3], px[0], grid._mesh, np.einsum("ij,ji->i", i_u.phi2D_lin(eta, 0.0), py) + py[3], py[0], px[3], px[0], grid._mesh, np.einsum("ij,ji->i", i_u.phi2D_lin(eta, 0.0), py), grid.deg2m ) def _create_selection_dict(dims, zdir=False): @@ -282,8 +282,8 @@ def _compute_corner_data(data, selection_dict) -> np.ndarray: V1 = corner_data[1, :] * c3 Vvel = (1 - eta) * V0 + eta * V1 - if grid._mesh == "spherical": - jac = i_u._compute_jacobian_determinant(py, px, eta, xsi) * 1852 * 60.0 + if grid._mesh.is_spherical(): + jac = i_u._compute_jacobian_determinant(py, px, eta, xsi) * grid.deg2m else: jac = i_u._compute_jacobian_determinant(py, px, eta, xsi) @@ -303,8 +303,8 @@ def _compute_corner_data(data, selection_dict) -> np.ndarray: u = u.compute() v = v.compute() - if grid._mesh == "spherical": - conversion = 1852 * 60.0 * np.cos(np.deg2rad(particle_positions["y"])) + if grid._mesh.is_spherical(): + conversion = grid.deg2m * np.cos(np.deg2rad(particle_positions["y"])) u /= conversion v /= conversion diff --git a/src/parcels/kernels/_advection.py b/src/parcels/kernels/_advection.py index 05df14c39..77162aad5 100644 --- a/src/parcels/kernels/_advection.py +++ b/src/parcels/kernels/_advection.py @@ -223,12 +223,12 @@ def AdvectionAnalytical(particles, fieldset): # pragma: no cover else: dz = 1.0 - c1 = i_u._geodetic_distance(py[0], py[1], px[0], px[1], grid.mesh, np.dot(i_u.phi2D_lin(0.0, xsi), py)) - c2 = i_u._geodetic_distance(py[1], py[2], px[1], px[2], grid.mesh, np.dot(i_u.phi2D_lin(eta, 1.0), py)) - c3 = i_u._geodetic_distance(py[2], py[3], px[2], px[3], grid.mesh, np.dot(i_u.phi2D_lin(1.0, xsi), py)) - c4 = i_u._geodetic_distance(py[3], py[0], px[3], px[0], grid.mesh, np.dot(i_u.phi2D_lin(eta, 0.0), py)) + c1 = i_u._geodetic_distance(py[0], py[1], px[0], px[1], grid.mesh, np.dot(i_u.phi2D_lin(0.0, xsi), py), grid.deg2m) + c2 = i_u._geodetic_distance(py[1], py[2], px[1], px[2], grid.mesh, np.dot(i_u.phi2D_lin(eta, 1.0), py), grid.deg2m) + c3 = i_u._geodetic_distance(py[2], py[3], px[2], px[3], grid.mesh, np.dot(i_u.phi2D_lin(1.0, xsi), py), grid.deg2m) + c4 = i_u._geodetic_distance(py[3], py[0], px[3], px[0], grid.mesh, np.dot(i_u.phi2D_lin(eta, 0.0), py), grid.deg2m) rad = np.pi / 180.0 - deg2m = 1852 * 60.0 + deg2m = grid.deg2m meshJac = (deg2m * deg2m * math.cos(rad * particles.y)) if grid.mesh == "spherical" else 1 dxdy = i_u._compute_jacobian_determinant(py, px, eta, xsi) * meshJac diff --git a/src/parcels/kernels/_advectiondiffusion.py b/src/parcels/kernels/_advectiondiffusion.py index 3593e2538..b5d327de9 100644 --- a/src/parcels/kernels/_advectiondiffusion.py +++ b/src/parcels/kernels/_advectiondiffusion.py @@ -8,14 +8,14 @@ __all__ = ["AdvectionDiffusionEM", "AdvectionDiffusionM1", "DiffusionUniformKh"] -def meters_to_degrees_zonal(deg, lat): # pragma: no cover +def meters_to_degrees_zonal(deg, lat, deg2m): # pragma: no cover """Convert square meters to square degrees longitude at a given latitude.""" - return deg / pow(1852 * 60.0 * np.cos(lat * np.pi / 180), 2) + return deg / pow(deg2m * np.cos(lat * np.pi / 180), 2) -def meters_to_degrees_meridional(deg): # pragma: no cover +def meters_to_degrees_meridional(deg, deg2m): # pragma: no cover """Convert square meters to square degrees latitude.""" - return deg / pow(1852 * 60.0, 2) + return deg / pow(deg2m, 2) def AdvectionDiffusionM1(particles, fieldset): # pragma: no cover @@ -39,27 +39,27 @@ def AdvectionDiffusionM1(particles, fieldset): # pragma: no cover Kxp1 = fieldset.Kh_zonal[particles.t, particles.z, particles.y, particles.x + fieldset.dres, particles] Kxm1 = fieldset.Kh_zonal[particles.t, particles.z, particles.y, particles.x - fieldset.dres, particles] - if fieldset.Kh_zonal.grid._mesh == "spherical": - Kxp1 = meters_to_degrees_zonal(Kxp1, particles.y) - Kxm1 = meters_to_degrees_zonal(Kxm1, particles.y) + if fieldset.Kh_zonal.grid._mesh.is_spherical(): + Kxp1 = meters_to_degrees_zonal(Kxp1, particles.y, fieldset.Kh_zonal.grid.deg2m) + Kxm1 = meters_to_degrees_zonal(Kxm1, particles.y, fieldset.Kh_zonal.grid.deg2m) dKdx = (Kxp1 - Kxm1) / (2 * fieldset.dres) u, v = fieldset.UV[particles.t, particles.z, particles.y, particles.x, particles] kh_zonal = fieldset.Kh_zonal[particles.t, particles.z, particles.y, particles.x, particles] - if fieldset.Kh_zonal.grid._mesh == "spherical": - kh_zonal = meters_to_degrees_zonal(kh_zonal, particles.y) + if fieldset.Kh_zonal.grid._mesh.is_spherical(): + kh_zonal = meters_to_degrees_zonal(kh_zonal, particles.y, fieldset.Kh_zonal.grid.deg2m) bx = np.sqrt(2 * kh_zonal) Kyp1 = fieldset.Kh_meridional[particles.t, particles.z, particles.y + fieldset.dres, particles.x, particles] Kym1 = fieldset.Kh_meridional[particles.t, particles.z, particles.y - fieldset.dres, particles.x, particles] - if fieldset.Kh_meridional.grid._mesh == "spherical": - Kyp1 = meters_to_degrees_meridional(Kyp1) - Kym1 = meters_to_degrees_meridional(Kym1) + if fieldset.Kh_meridional.grid._mesh.is_spherical(): + Kyp1 = meters_to_degrees_meridional(Kyp1, fieldset.Kh_meridional.grid.deg2m) + Kym1 = meters_to_degrees_meridional(Kym1, fieldset.Kh_meridional.grid.deg2m) dKdy = (Kyp1 - Kym1) / (2 * fieldset.dres) kh_meridional = fieldset.Kh_meridional[particles.t, particles.z, particles.y, particles.x, particles] - if fieldset.Kh_meridional.grid._mesh == "spherical": - kh_meridional = meters_to_degrees_meridional(kh_meridional) + if fieldset.Kh_meridional.grid._mesh.is_spherical(): + kh_meridional = meters_to_degrees_meridional(kh_meridional, fieldset.Kh_meridional.grid.deg2m) by = np.sqrt(2 * kh_meridional) # Particle positions are updated only after evaluating all terms. @@ -88,28 +88,28 @@ def AdvectionDiffusionEM(particles, fieldset): # pragma: no cover Kxp1 = fieldset.Kh_zonal[particles.t, particles.z, particles.y, particles.x + fieldset.dres, particles] Kxm1 = fieldset.Kh_zonal[particles.t, particles.z, particles.y, particles.x - fieldset.dres, particles] - if fieldset.Kh_zonal.grid._mesh == "spherical": - Kxp1 = meters_to_degrees_zonal(Kxp1, particles.y) - Kxm1 = meters_to_degrees_zonal(Kxm1, particles.y) + if fieldset.Kh_zonal.grid._mesh.is_spherical(): + Kxp1 = meters_to_degrees_zonal(Kxp1, particles.y, fieldset.Kh_zonal.grid.deg2m) + Kxm1 = meters_to_degrees_zonal(Kxm1, particles.y, fieldset.Kh_zonal.grid.deg2m) dKdx = (Kxp1 - Kxm1) / (2 * fieldset.dres) ax = u + dKdx kh_zonal = fieldset.Kh_zonal[particles.t, particles.z, particles.y, particles.x, particles] - if fieldset.Kh_zonal.grid._mesh == "spherical": - kh_zonal = meters_to_degrees_zonal(kh_zonal, particles.y) + if fieldset.Kh_zonal.grid._mesh.is_spherical(): + kh_zonal = meters_to_degrees_zonal(kh_zonal, particles.y, fieldset.Kh_zonal.grid.deg2m) bx = np.sqrt(2 * kh_zonal) Kyp1 = fieldset.Kh_meridional[particles.t, particles.z, particles.y + fieldset.dres, particles.x, particles] Kym1 = fieldset.Kh_meridional[particles.t, particles.z, particles.y - fieldset.dres, particles.x, particles] - if fieldset.Kh_meridional.grid._mesh == "spherical": - Kyp1 = meters_to_degrees_meridional(Kyp1) - Kym1 = meters_to_degrees_meridional(Kym1) + if fieldset.Kh_meridional.grid._mesh.is_spherical(): + Kyp1 = meters_to_degrees_meridional(Kyp1, fieldset.Kh_meridional.grid.deg2m) + Kym1 = meters_to_degrees_meridional(Kym1, fieldset.Kh_meridional.grid.deg2m) dKdy = (Kyp1 - Kym1) / (2 * fieldset.dres) ay = v + dKdy kh_meridional = fieldset.Kh_meridional[particles.t, particles.z, particles.y, particles.x, particles] - if fieldset.Kh_meridional.grid._mesh == "spherical": - kh_meridional = meters_to_degrees_meridional(kh_meridional) + if fieldset.Kh_meridional.grid._mesh.is_spherical(): + kh_meridional = meters_to_degrees_meridional(kh_meridional, fieldset.Kh_meridional.grid.deg2m) by = np.sqrt(2 * kh_meridional) # Particle positions are updated only after evaluating all terms. @@ -142,9 +142,9 @@ def DiffusionUniformKh(particles, fieldset): # pragma: no cover kh_zonal = fieldset.Kh_zonal[particles] kh_meridional = fieldset.Kh_meridional[particles] - if fieldset.Kh_zonal.grid._mesh == "spherical": - kh_zonal = meters_to_degrees_zonal(kh_zonal, particles.y) - kh_meridional = meters_to_degrees_meridional(kh_meridional) + if fieldset.Kh_zonal.grid._mesh.is_spherical(): + kh_zonal = meters_to_degrees_zonal(kh_zonal, particles.y, fieldset.Kh_zonal.grid.deg2m) + kh_meridional = meters_to_degrees_meridional(kh_meridional, fieldset.Kh_meridional.grid.deg2m) bx = np.sqrt(2 * kh_zonal) by = np.sqrt(2 * kh_meridional) diff --git a/tests/test_fieldset.py b/tests/test_fieldset.py index abdd26062..995021c18 100644 --- a/tests/test_fieldset.py +++ b/tests/test_fieldset.py @@ -454,8 +454,8 @@ def test_fieldset_describe_backends(tmp_path): | UV | VectorField | 0 | CGrid_Velocity(...) | - | | UVW | VectorField | 0 | CGrid_Velocity(...) | - | -mesh: spherical -time interval: (np.datetime64('2000-01-02T12:00:00.000000000'), np.datetime64('2000-01-12T12:00:00.000000000')) +mesh: SphericalMesh(radius=6366707.019493707) +time interval: (np.datetime64('2000-01-02T12:00:00.000000000'), np.datetime64('2000-01-27T12:00:00.000000000')) """ fieldset.describe(io) actual = io.getvalue() @@ -474,8 +474,8 @@ def test_fieldset_describe_backends(tmp_path): | UV | VectorField | 0 | CGrid_Velocity(...) | - | | UVW | VectorField | 0 | CGrid_Velocity(...) | - | -mesh: spherical -time interval: (np.datetime64('2000-01-02T12:00:00.000000000'), np.datetime64('2000-01-12T12:00:00.000000000')) +mesh: SphericalMesh(radius=6366707.019493707) +time interval: (np.datetime64('2000-01-02T12:00:00.000000000'), np.datetime64('2000-01-27T12:00:00.000000000')) """ fieldset.describe(io) actual = io.getvalue() @@ -496,8 +496,8 @@ def test_fieldset_describe_backends(tmp_path): | UV | VectorField | 0 | CGrid_Velocity(...) | - | | UVW | VectorField | 0 | CGrid_Velocity(...) | - | -mesh: spherical -time interval: (np.datetime64('2000-01-02T12:00:00.000000000'), np.datetime64('2000-01-12T12:00:00.000000000')) +mesh: SphericalMesh(radius=6366707.019493707) +time interval: (np.datetime64('2000-01-02T12:00:00.000000000'), np.datetime64('2000-01-27T12:00:00.000000000')) """ fieldset.describe(io) actual = io.getvalue() diff --git a/tests/test_mesh.py b/tests/test_mesh.py new file mode 100644 index 000000000..9a0a4b709 --- /dev/null +++ b/tests/test_mesh.py @@ -0,0 +1,84 @@ +import numpy as np +import pytest + +from parcels import FieldSet, ParticleSet, SphericalMesh +from parcels._core.mesh import EARTH_RADIUS +from parcels._datasets.structured.generated import simple_UV_dataset +from parcels.kernels import AdvectionRK4 + +EARTH_DEG2M = EARTH_RADIUS * np.pi / 180 + + +def test_spherical_mesh_deg2m(): + assert SphericalMesh().radius == EARTH_RADIUS + assert SphericalMesh().deg2m == EARTH_DEG2M + r = 3389500.0 # Mars radius + assert SphericalMesh(radius=r).deg2m == pytest.approx(r * np.pi / 180) + + +@pytest.mark.parametrize( + "mesh, exp_radius, exp_deg2m", + [ + ("spherical", EARTH_RADIUS, EARTH_DEG2M), + (SphericalMesh(), EARTH_RADIUS, EARTH_DEG2M), + (SphericalMesh(radius=3389500.0), 3389500.0, 3389500.0 * np.pi / 180), + ], +) +def test_xgrid_radius_and_deg2m(mesh, exp_radius, exp_deg2m): + grid = FieldSet.from_sgrid_conventions(simple_UV_dataset(), mesh=mesh).U.grid + assert isinstance(grid._mesh, SphericalMesh) + assert grid._mesh.radius == exp_radius + assert grid.deg2m == pytest.approx(exp_deg2m) + + +@pytest.mark.parametrize( + "mesh, deg2m", + [ + (SphericalMesh(), EARTH_DEG2M), # No radius entered + (SphericalMesh(radius=3389500.0), 3389500.0 * np.pi / 180), # Mars + (SphericalMesh(radius=6051800.0), 6051800.0 * np.pi / 180), # Venus + (SphericalMesh(radius=EARTH_RADIUS), EARTH_DEG2M), # Explicit Earth + ], +) +def test_advection_uses_custom_radius(mesh, deg2m, npart=10): + ds = simple_UV_dataset() + ds["U"].data[:] = 1.0 + fieldset = FieldSet.from_sgrid_conventions(ds, mesh=mesh) + + runtime = 7200 + startlat = np.linspace(0, 80, npart) + startlon = 20.0 + np.zeros(npart) + pset = ParticleSet(fieldset, x=startlon, y=startlat) + pset.execute(AdvectionRK4, runtime=runtime, dt=np.timedelta64(15, "m")) + + expected_dlon = runtime / (deg2m * np.cos(np.deg2rad(pset.y))) + np.testing.assert_allclose(pset.x - startlon, expected_dlon, atol=1e-5) + np.testing.assert_allclose(pset.y, startlat, atol=1e-5) + + +def test_advection_flat_mesh(npart=10): + ds = simple_UV_dataset(mesh="flat") + ds["U"].data[:] = 1.0 + fieldset = FieldSet.from_sgrid_conventions(ds, mesh="flat") + + runtime = 7200 + startlat = np.linspace(0, 80, npart) + startlon = 20.0 + np.zeros(npart) + pset = ParticleSet(fieldset, x=startlon, y=startlat) + pset.execute(AdvectionRK4, runtime=runtime, dt=np.timedelta64(15, "m")) + + assert fieldset.U.grid.deg2m == 1.0 # flat mesh deg2m + np.testing.assert_allclose(pset.x - startlon, runtime, atol=1e-5) + np.testing.assert_allclose(pset.y, startlat, atol=1e-5) + + +@pytest.mark.parametrize("bad_radius", ["6371000", [6371000], (1, 2), {}]) +def test_spherical_mesh_rejects_non_numeric_radius(bad_radius): + with pytest.raises(TypeError): + SphericalMesh(radius=bad_radius) + + +@pytest.mark.parametrize("bad_radius", [0, -1.0, -6371000]) +def test_spherical_mesh_rejects_nonpos_radius(bad_radius): + with pytest.raises(ValueError): + SphericalMesh(radius=bad_radius) diff --git a/tests/test_typing.py b/tests/test_typing.py index 4f1d3e0c5..e69de29bb 100644 --- a/tests/test_typing.py +++ b/tests/test_typing.py @@ -1,20 +0,0 @@ -import numpy as np -import pytest -import xarray as xr - -from parcels._typing import ( - assert_valid_mesh, -) - - -def test_invalid_assert_valid_mesh(): - with pytest.raises(ValueError, match="Invalid value"): - assert_valid_mesh("invalid option") - - ds = xr.Dataset({"A": (("a", "b"), np.arange(20).reshape(4, 5))}) - with pytest.raises(ValueError, match="Invalid input type"): - assert_valid_mesh(ds) - - -def test_assert_valid_mesh(): - assert_valid_mesh("spherical") diff --git a/tests/test_uxgrid.py b/tests/test_uxgrid.py index 5e3038cf3..63df6a3ca 100644 --- a/tests/test_uxgrid.py +++ b/tests/test_uxgrid.py @@ -22,7 +22,8 @@ def test_uxgrid_axes(uxds): @pytest.mark.parametrize("mesh", ["flat", "spherical"]) def test_uxgrid_mesh(uxds, mesh): grid = UxGrid(uxds.uxgrid, z=uxds.coords["zf"], mesh=mesh) - assert grid._mesh == mesh + + assert mesh in grid._mesh.__class__.__name__.lower() @pytest.mark.parametrize("uxds", [uxdatasets["stommel_gyre_delaunay"]]) diff --git a/tests/test_xgrid.py b/tests/test_xgrid.py index 2ba04e115..d957ab0b9 100644 --- a/tests/test_xgrid.py +++ b/tests/test_xgrid.py @@ -72,7 +72,7 @@ def test_xgrid_axes(fieldset): @pytest.mark.parametrize("mesh", ["flat", "spherical"]) def test_uxgrid_mesh(ds, mesh): grid = FieldSet.from_sgrid_conventions(ds, mesh=mesh).data_g.grid - assert grid._mesh == mesh + assert mesh in grid._mesh.__class__.__name__.lower() @pytest.mark.skip(