Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
21 commits
Select commit Hold shift + click to select a range
ba7fe37
Initial 3d z coordinate implementation for unstructured grids
wyatt-fluidnumerics Oct 5, 2026
57cf6dd
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] Oct 5, 2026
4f921d6
Removing unused polars import
erikvansebille Oct 6, 2026
a12142d
Updated markdown cells in new tutorial
wyatt-fluidnumerics Oct 6, 2026
cf3725b
Merge branch 'sigma-vertical-grid-support' of github.com:Parcels-code…
wyatt-fluidnumerics Oct 6, 2026
85aa9d4
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] Oct 6, 2026
cc71804
Merge branch 'main' into sigma-vertical-grid-support
wyatt-fluidnumerics Oct 7, 2026
21c2cf5
Temporal Linear interpolatation of the z coordinate during grid search
wyatt-fluidnumerics Oct 8, 2026
c9bc913
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] Oct 8, 2026
37ac028
Merge branch 'main' into sigma-vertical-grid-support
wyatt-fluidnumerics Oct 8, 2026
dea15dc
Fixed ci failures
wyatt-fluidnumerics Oct 8, 2026
debe984
Added experimental/memory warning during 3D z coordinate fieldset con…
wyatt-fluidnumerics Oct 8, 2026
85c4d69
Merge branch 'main' into sigma-vertical-grid-support
wyatt-fluidnumerics Oct 9, 2026
09fee29
time varying z coordinate support for XGrids
wyatt-fluidnumerics Oct 9, 2026
0701f1c
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] Oct 9, 2026
0ab3022
Changed tutorial to include structured grids
wyatt-fluidnumerics Oct 9, 2026
a6788a0
Merge branch 'sigma-vertical-grid-support' of github.com:Parcels-code…
wyatt-fluidnumerics Oct 9, 2026
27b8192
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] Oct 9, 2026
096033d
Enforce correct dimensions for Z on XGrid construction
wyatt-fluidnumerics Oct 9, 2026
6dfa56d
Merge branch 'sigma-vertical-grid-support' of github.com:Parcels-code…
wyatt-fluidnumerics Oct 9, 2026
3ce43e4
Merge branch 'main' into sigma-vertical-grid-support
wyatt-fluidnumerics Oct 9, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
504 changes: 504 additions & 0 deletions docs/user_guide/examples/tutorial_sigma_coordinates.ipynb

Large diffs are not rendered by default.

1 change: 1 addition & 0 deletions docs/user_guide/index.md
Original file line number Diff line number Diff line change
Expand Up @@ -51,6 +51,7 @@ examples/tutorial_schism.ipynb
:name: work-with-fieldsets
:titlesonly:
examples/explanation_grids.ipynb
examples/tutorial_sigma_coordinates.ipynb
examples/tutorial_velocityconversion.ipynb
examples/tutorial_nestedgrids.ipynb
examples/tutorial_manipulating_field_data.ipynb
Expand Down
8 changes: 7 additions & 1 deletion src/parcels/_core/basegrid.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,7 +29,9 @@ class BaseGrid(ABC):
_mesh: FlatMesh | SphericalMesh

@abstractmethod
def search(self, z: float, y: float, x: float, ei=None) -> dict[str, tuple[int, float | np.ndarray]]:
def search(
self, z: float, y: float, x: float, ei=None, ti=None, tau=None
) -> dict[str, tuple[int, float | np.ndarray]]:
"""
Perform a spatial (and optionally vertical) search to locate the grid element
that contains a given point (x, y, z).
Expand All @@ -49,6 +51,10 @@ def search(self, z: float, y: float, x: float, ei=None) -> dict[str, tuple[int,
A previously computed encoded index (e.g., raveled face or cell index). If provided,
the search will first attempt to validate and reuse it before falling back to
a global or local search strategy.
ti : np.ndarray, optional
Time index of each query point, as returned by ``_search_time_index``.
tau : np.ndarray, optional
Barycentric time coordinate of each query point, as returned by ``_search_time_index``.
search2D : bool, default=False
If True, perform only a 2D search (x, y), ignoring the vertical component z.

Expand Down
7 changes: 5 additions & 2 deletions src/parcels/_core/field.py
Original file line number Diff line number Diff line change
Expand Up @@ -398,8 +398,11 @@ def _get_positions(field: Field, t, z, y, x, particles, _ei) -> tuple[dict, dict
raise ValueError(f"Time values for particles with indices {nan_indices} cannot be NaN.")
particle_positions = {"t": t, "z": z, "y": y, "x": x}
grid_positions = {}
grid_positions.update(_search_time_index(field, t))
grid_positions.update(field.grid.search(z, y, x, ei=_ei))
time_positions = _search_time_index(field, t)
grid_positions.update(time_positions)
grid_positions.update(
field.grid.search(z, y, x, ei=_ei, ti=time_positions["T"]["index"], tau=time_positions["T"]["bcoord"])
)
_update_particles_ei(particles, grid_positions, field)
_update_particle_states_position(particles, grid_positions)
return particle_positions, grid_positions
39 changes: 39 additions & 0 deletions src/parcels/_core/index_search.py
Original file line number Diff line number Diff line change
Expand Up @@ -64,6 +64,45 @@ def _search_1d_array(
return np.atleast_1d(index), np.atleast_1d(bcoord)


def _search_1d_columns(columns: np.ndarray, x: np.ndarray) -> tuple[np.ndarray, np.ndarray]:
"""
Searches for particle locations in per-particle 1D columns and returns barycentric coordinate along dimension.

Row-wise counterpart of ``_search_1d_array``: particle p searches only its own column ``columns[p]``.

Assumptions:
- each column is strictly monotonically increasing.

Parameters
----------
columns : np.ndarray
2D array of shape (n_particles, n_levels), one column per particle.
x : np.ndarray
Position of each particle along its column, shape (n_particles,).

Returns
-------
array of int
Index of the element just before the position x in each column. Note that this index is -2 if the index is left out of bounds and -1 if the index is right out of bounds.
array of float
Barycentric coordinate.
"""
n_levels = columns.shape[1]
if n_levels < 2:
return np.zeros(shape=x.shape, dtype=np.int32), np.zeros_like(x)
# The number of column entries strictly below x equals np.searchsorted(column, x, side="left")
index = np.clip((columns < x[:, None]).sum(axis=1) - 1, 0, n_levels - 2)
rows = np.arange(columns.shape[0])
lower = columns[rows, index]
upper = columns[rows, index + 1]
bcoord = (x - lower) / (upper - lower)

index = np.where(x < columns[:, 0], LEFT_OUT_OF_BOUNDS, index)
index = np.where(x > columns[:, -1], RIGHT_OUT_OF_BOUNDS, index)

return index, bcoord


def _search_time_index(field: Field, time: np.ndarray):
"""Find and return the index and relative coordinate in the time array associated with a given time.

Expand Down
3 changes: 3 additions & 0 deletions src/parcels/_core/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -359,6 +359,9 @@ def __init__(self, data: ux.UxDataset, grid: UxGrid, vector_field_components: pt
if not isinstance(grid, UxGrid):
raise ValueError(f"Expected `grid` to be a Parcels UxGrid object. Got {type(grid)}.")

if grid.z.ndim == 3 and not grid.z["time"].equals(data["time"]):
raise ValueError("A time-varying (3D) z must have the same time coordinate as `data`.")

self.data = data
self.grid = grid
self.vector_field_components = vector_field_components
Expand Down
10 changes: 10 additions & 0 deletions src/parcels/_core/particleset.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,9 @@
float_to_datelike,
timedelta_to_float,
)
from parcels._core.uxgrid import UxGrid
from parcels._core.warnings import ParticleSetWarning
from parcels._core.xgrid import XGrid
from parcels._logger import logger

__all__ = ["ParticleSet"]
Expand Down Expand Up @@ -82,6 +84,14 @@ def __init__(
if z is None:
minz = None
for field in self.fieldset.fields.values():
has_time_varying_z = (isinstance(field.grid, UxGrid) and field.grid.z.ndim == 3) or (
isinstance(field.grid, XGrid) and "Z" in field.grid.axes and field.grid._ds["depth"].ndim == 4
)
if has_time_varying_z:
raise ValueError(
f"Field {field.name!r} has a time-varying vertical grid, so there is no default "
"particle depth. Pass the particle depths explicitly with `z`."
)
for depth in field.grid.depth:
if minz is None or np.abs(depth) < np.abs(minz):
minz = depth
Expand Down
78 changes: 69 additions & 9 deletions src/parcels/_core/uxgrid.py
Original file line number Diff line number Diff line change
@@ -1,13 +1,16 @@
from __future__ import annotations

import warnings
from typing import TYPE_CHECKING, Literal

import numpy as np
import xarray as xr
from dask import is_dask_collection

from parcels._core.basegrid import BaseGrid
from parcels._core.index_search import GRID_SEARCH_ERROR, _search_1d_array, uxgrid_point_in_cell
from parcels._core.index_search import GRID_SEARCH_ERROR, _search_1d_columns, uxgrid_point_in_cell
from parcels._core.mesh import SphericalMesh, get_mesh
from parcels._core.warnings import FieldSetWarning

if TYPE_CHECKING:
import uxarray as ux
Expand All @@ -32,9 +35,9 @@ def __init__(
grid : ux.grid.Grid
The uxarray grid object containing the unstructured grid data.
z : ux.UxDataArray
A 1D array of vertical coordinates (depths) associated with the layer interface heights (not the mid-layer depths).
While uxarray allows nz to be spatially and temporally varying, the parcels.UxGrid class considers the case where
the vertical coordinate is constant in time and space. This implies flat bottom topography and no moving ALE vertical grid.
Vertical coordinates (depths) of the layer interface heights (not the mid-layer depths). Either a 1D array,
constant in time and space (flat bottom topography, no moving vertical grid), or a 3D array with dims
("time", "zf", "n_node"), varying in time and across the mesh nodes.
mesh : str
The type of mesh used for the grid. Either "flat" or "spherical".
"""
Expand All @@ -45,8 +48,20 @@ def __init__(
self.uxgrid = grid
if not isinstance(z, ux.UxDataArray):
raise TypeError("z must be an instance of ux.UxDataArray")
if z.ndim != 1:
raise ValueError("z must be a 1D array of vertical coordinates")
if z.ndim not in (1, 3):
raise ValueError(f"z must be a 1D or 3D array of vertical coordinates, got {z.ndim}D")
if z.ndim == 3 and z.dims != ("time", "zf", "n_node"):
raise ValueError(f"A 3D z must have dims ('time', 'zf', 'n_node'), got {z.dims}")
if z.ndim == 3:
warnings.warn(
"Time-varying (3D) z coordinates are experimental and may cause significant memory overhead that "
f"leads to OOM errors. This z coordinate has sizes {dict(z.sizes)} ({z.nbytes / 1e9:.3g} GB). "
"Assumptions: z is defined at the layer interfaces ('zf') on the mesh nodes ('n_node') and is strictly "
"increasing along 'zf'; each particle's z column is interpolated barycentrically from its face's "
"nodes and linearly in time between z snapshots.",
FieldSetWarning,
stacklevel=4,
)
self.z = z
self._mesh = get_mesh(mesh)
self._spatialhash = None
Expand Down Expand Up @@ -76,6 +91,8 @@ def get_axis_dim(self, axis: _UXGRID_AXES) -> int:
raise ValueError(f"Axis {axis!r} is not part of this grid. Available axes: {self.axes}")

if axis == "Z":
if self.z.ndim == 3:
return self.z.sizes["zf"]
return len(self.z.values)
elif axis == "FACE":
return self.uxgrid.n_face
Expand All @@ -87,7 +104,7 @@ def deg2m(self) -> float:
return self._mesh.deg2m
return 1.0

def search(self, z, y, x, ei=None, tol=1e-6):
def search(self, z, y, x, ei=None, ti=None, tau=None, tol=1e-6):
"""
Search for the grid cell (face) and vertical layer that contains the given points.

Expand All @@ -105,15 +122,19 @@ def search(self, z, y, x, ei=None, tol=1e-6):
TO BE IMPLEMENTED : If provided, we'll check
if the points are within the faces specified by these indices. For cells where the particles
are not found, a nearest neighbor search will be performed. As a last resort, the spatial hash will be used.
ti : np.ndarray, optional
Time index of each point, as returned by ``_search_time_index``. Required when z is 3D; selects the
earlier of the two z snapshots (``ti`` and ``ti + 1``) that are linearly interpolated with ``tau``.
tau : np.ndarray, optional
Barycentric time coordinate of each point, as returned by ``_search_time_index``. Required when z is 3D;
used for linear interpolation of the z coordinate for the construction of a particle's column.
tol : float, optional
Tolerance for barycentric coordinate checks. Default is 1e-6.
"""
x = np.asarray(x, dtype=np.float32)
y = np.asarray(y, dtype=np.float32)
z = np.asarray(z, dtype=np.float32)

zi, zeta = _search_1d_array(self.z.values, z)

if np.any(ei):
indices = self.unravel_index(ei)
fi = indices.get("FACE")
Expand All @@ -134,4 +155,43 @@ def search(self, z, y, x, ei=None, tol=1e-6):
coords[zero_indices, :] = coords_q
fi[zero_indices] = face_ids_q

found = fi >= 0
if self.z.ndim == 3:
if ti is None or tau is None:
raise ValueError(
"Searching a UxGrid with a time-varying (3D) z requires the time index ti and barycentric coordinate tau"
)

cols_ti = self.z.isel(
time=xr.DataArray(np.broadcast_to(ti, fi.shape)[found], dims="points"),
n_node=xr.DataArray(self.uxgrid.face_node_connectivity[fi[found], :].values, dims=("points", "nodes")),
ignore_grid=True,
).compute()

if self.z.shape[0] == 1:
node_columns = cols_ti
else:
cols_tnext = self.z.isel(
time=xr.DataArray(np.broadcast_to(ti + 1, fi.shape)[found], dims="points"),
n_node=xr.DataArray(
self.uxgrid.face_node_connectivity[fi[found], :].values, dims=("points", "nodes")
),
ignore_grid=True,
).compute()

tau_points = xr.DataArray(np.broadcast_to(tau, fi.shape)[found], dims="points")
node_columns = (1 - tau_points) * cols_ti + tau_points * cols_tnext

bcoords = xr.DataArray(coords[found], dims=("points", "nodes"))
# Particles outside the mesh are given a NaN z column for vertical searching
columns = np.full((fi.size, self.z.sizes["zf"]), np.nan)
columns[found] = xr.dot(node_columns, bcoords, dim="nodes").transpose("points", "zf").values
else:
columns = np.broadcast_to(self.z.values, (z.size, self.z.size))

zi, zeta = _search_1d_columns(columns, z)
# Particles outside the mesh are given a 0 vertical index and a NaN vertical barycentric coordinate
zi = np.where(found, zi, 0)
zeta = np.where(found, zeta, np.nan)

return {"Z": {"index": zi, "bcoord": zeta}, "FACE": {"index": fi, "bcoord": coords}}
Loading
Loading