Source code for stilt.config.spatial

"""Bounding boxes and footprint grids."""

from __future__ import annotations

from typing import TYPE_CHECKING, Literal, cast

import numpy as np
from pydantic import BaseModel, ConfigDict, Field

if TYPE_CHECKING:
    import pandas as pd
    import xarray as xr

VerticalReference = Literal["agl", "msl"]


def validate_vertical_reference(reference: str) -> VerticalReference:
    """Return ``reference`` in lower case, raising unless it is ``agl`` or ``msl``."""
    normalized = reference.lower()
    if normalized not in {"agl", "msl"}:
        raise ValueError(
            f"Vertical reference must be 'agl' or 'msl'. Got {reference!r}."
        )
    return cast(VerticalReference, normalized)


def kmsl_from_vertical_reference(reference: VerticalReference) -> int:
    """Return HYSPLIT's ``KMSL`` value for a vertical reference: 0 for AGL, 1 for MSL."""
    return 0 if reference == "agl" else 1


def _grid_cell_starts(minimum: float, maximum: float, resolution: float) -> np.ndarray:
    """
    Return the lower edges of the whole cells between two grid bounds.

    Cells start at ``minimum`` and step by ``resolution`` while the whole cell
    fits inside ``[minimum, maximum]``, so ``maximum`` is the outer edge of
    the last cell.

    Decimal bounds are not exact in binary, so ``maximum - minimum`` carries a
    rounding error that grows with the bounds (``40.93 - 40.45`` is
    ``0.4799999999999969``). The tolerance is sized to that error so the last
    intended cell is kept (48 cells here at 0.01, as STILT-R's ``seq()``
    gives) and a partial cell is still dropped.
    """
    if resolution <= 0:
        raise ValueError("Grid resolution must be positive.")
    quotient = (maximum - minimum) / resolution
    scale = max(abs(minimum), abs(maximum), abs(maximum - minimum))
    tol = 16 * np.finfo(float).eps * scale / resolution
    n_cells = int(np.floor(quotient + tol))
    if n_cells < 1:
        raise ValueError("Grid extent must contain at least one complete cell.")
    return minimum + np.arange(n_cells, dtype=float) * resolution


def _cf_grid_mapping_attrs(projection: str) -> dict[str, object]:
    """Return CF-style grid-mapping attributes for a PROJ string."""
    attrs: dict[str, object] = {"proj4_params": projection}
    from pyproj import CRS

    crs = CRS.from_user_input(projection)
    attrs.update(crs.to_cf())
    wkt = crs.to_wkt()
    attrs["spatial_ref"] = wkt
    attrs["crs_wkt"] = wkt
    return {
        key: value
        for key, value in attrs.items()
        if isinstance(value, str | int | float | np.number)
    }


def cf_axis_attrs(dim: str) -> dict[str, str]:
    """
    Return the CF attributes for one coordinate axis: ``lon``, ``lat``, ``x`` or ``y``.

    Grids and footprints both use this, so their axes carry the same
    attributes.
    """
    return {
        "lon": {
            "standard_name": "longitude",
            "long_name": "longitude",
            "units": "degrees_east",
            "axis": "X",
        },
        "lat": {
            "standard_name": "latitude",
            "long_name": "latitude",
            "units": "degrees_north",
            "axis": "Y",
        },
        "x": {
            "standard_name": "projection_x_coordinate",
            "long_name": "x coordinate of projection",
            "units": "m",
            "axis": "X",
        },
        "y": {
            "standard_name": "projection_y_coordinate",
            "long_name": "y coordinate of projection",
            "units": "m",
            "axis": "Y",
        },
    }[dim]


[docs] class Bounds(BaseModel): """Longitude/latitude bounding box, in degrees.""" model_config = ConfigDict(frozen=True) xmin: float = Field(..., description="Western edge, in degrees longitude.") xmax: float = Field(..., description="Eastern edge, in degrees longitude.") ymin: float = Field(..., description="Southern edge, in degrees latitude.") ymax: float = Field(..., description="Northern edge, in degrees latitude.")
[docs] class Grid(Bounds): """ Footprint grid: longitude/latitude bounds, cell size, and projection. The bounds are always longitude/latitude. With a projected ``projection``, the grid covers the bounds' extent in that projection and ``xres`` and ``yres`` are in its units. """ model_config = ConfigDict(frozen=True) xres: float = Field( ..., description="Cell width in projection units (degrees for longlat, meters for UTM).", ) yres: float = Field( ..., description="Cell height in projection units (degrees for longlat, meters for UTM).", ) projection: str = Field( "+proj=longlat", description=( "Projection of the footprint grid, as a PROJ string. Particles and " "bounds are projected to it before gridding." ), ) @property def resolution(self) -> str: """Cell size as text, such as ``'0.01x0.01'``.""" return f"{self.xres}x{self.yres}" @property def is_longlat(self) -> bool: """Whether the grid is in longitude/latitude degrees.""" return "+proj=longlat" in self.projection @property def min_cell_width(self) -> float: """Smaller of ``xres`` and ``yres``, as other geometries report it.""" return float(min(self.xres, self.yres))
[docs] @classmethod def from_geometry( cls, geometry, *, cells_per_target: float = 4.0, projection: str | None = None, max_cells: int = 50_000_000, ) -> Grid: """ Return a grid fine enough to resolve the cells of a geometry. The bounds are the geometry's extent, rounded outward to whole cells. The cell size is the smallest geometry cell width divided by ``cells_per_target``, rounded down to one significant figure. Parameters ---------- geometry : Mesh, Zones, or Grid Geometry to resolve. Any object with ``bounds``, ``min_cell_width``, ``crs`` or ``projection``, and ``is_longlat``. cells_per_target : float, default 4 Grid cells across the smallest geometry cell. projection : str, optional Projection of the grid. Defaults to the geometry's CRS; the geometry is reprojected when they differ. max_cells : int, default 50_000_000 Warn when the grid would have more cells than this. Returns ------- Grid The derived grid, with longitude/latitude bounds. """ import math import warnings crs = getattr(geometry, "crs", None) or getattr(geometry, "projection", None) if not isinstance(crs, str): raise TypeError("geometry must expose a 'crs' or 'projection' string.") if projection is not None and projection != crs: from stilt.geometry import Mesh, Zones geometry = geometry.base if isinstance(geometry, Zones) else geometry if isinstance(geometry, Grid): geometry = Mesh.from_grid(geometry) geometry = geometry.to_crs(projection) crs = projection projection = crs width = float(geometry.min_cell_width) / float(cells_per_target) if width <= 0: raise ValueError("Geometry cells must have positive width.") exp = math.floor(math.log10(width)) res = math.floor(width / 10**exp) * 10**exp # round down, 1 sig fig res = float(f"{res:.1g}") xmin, ymin, xmax, ymax = geometry.bounds xmin, ymin = math.floor(xmin / res) * res, math.floor(ymin / res) * res xmax, ymax = math.ceil(xmax / res) * res, math.ceil(ymax / res) * res if xmax <= xmin: xmax = xmin + res if ymax <= ymin: ymax = ymin + res n_cells = ((xmax - xmin) / res) * ((ymax - ymin) / res) if n_cells > max_cells: warnings.warn( f"Derived grid has ~{n_cells:.3g} cells at resolution {res:g}; " "consider a coarser cells_per_target or a smaller domain.", stacklevel=2, ) if not geometry.is_longlat: # Bounds are always lon/lat: back-transform the snapped envelope. # Pad by one cell first; ``axes`` re-projects the lon/lat corners # and takes their extremes, which can shave an edge cell off a # rotated projection otherwise. xmin, xmax, ymin, ymax = xmin - res, xmax + res, ymin - res, ymax + res from pyproj import Transformer tr = Transformer.from_crs(projection, "EPSG:4326", always_xy=True) xs, ys = tr.transform([xmin, xmax, xmin, xmax], [ymin, ymin, ymax, ymax]) xmin, xmax = float(min(xs)), float(max(xs)) ymin, ymax = float(min(ys)), float(max(ys)) return cls( xmin=float(xmin), xmax=float(xmax), ymin=float(ymin), ymax=float(ymax), xres=res, yres=res, projection=projection, )
[docs] @classmethod def from_geometries(cls, geometries, **kwargs) -> Grid: """ Return one grid that resolves several geometries. The grid covers all their extents at the finest cell size any of them needs. Keyword arguments are passed to :meth:`from_geometry`. """ grids = [cls.from_geometry(g, **kwargs) for g in geometries] res = min(min(g.xres, g.yres) for g in grids) projection = grids[0].projection return cls( xmin=min(g.xmin for g in grids), xmax=max(g.xmax for g in grids), ymin=min(g.ymin for g in grids), ymax=max(g.ymax for g in grids), xres=res, yres=res, projection=projection, )
@property def axes(self) -> tuple[np.ndarray, np.ndarray]: """ Cell-center coordinates ``(x, y)`` in projection units, ascending. These are the footprint's coordinates on this grid. They are rounded to 10 decimals (``40.35`` rather than ``40.349999999999994``) so they match labels built elsewhere from the same bounds and resolution. Projected grids need ``pyproj``. """ xmin, xmax, ymin, ymax = self.xmin, self.xmax, self.ymin, self.ymax if not self.is_longlat: from pyproj import Transformer tr = Transformer.from_crs("EPSG:4326", self.projection, always_xy=True) corners_x, corners_y = tr.transform([xmin, xmax], [ymin, ymax]) xmin, xmax = float(np.min(corners_x)), float(np.max(corners_x)) ymin, ymax = float(np.min(corners_y)), float(np.max(corners_y)) x_centers = _grid_cell_starts(xmin, xmax, self.xres) + self.xres / 2 y_centers = _grid_cell_starts(ymin, ymax, self.yres) + self.yres / 2 return np.round(x_centers, 10), np.round(y_centers, 10) @property def cells(self) -> tuple[np.ndarray, np.ndarray]: """Every cell center as flat ``(x, y)`` arrays, with ``x`` varying slowest.""" x, y = self.axes xx, yy = np.meshgrid(x, y, indexing="ij") return xx.ravel(), yy.ravel() @property def index(self) -> pd.MultiIndex: """ Index of every cell, in the same order as ``cells``. The levels are named ``lon`` and ``lat`` for a longitude/latitude grid and ``x`` and ``y`` for a projected one, as in the footprint. """ import pandas as pd x, y = self.axes names = ["lon", "lat"] if self.is_longlat else ["x", "y"] return pd.MultiIndex.from_product([x, y], names=names)
[docs] def to_xarray(self) -> xr.Dataset: """ Return the grid as a CF-1.8 dataset of cell centers. The dataset has ``lon`` and ``lat`` coordinates (``x`` and ``y`` when projected) matching the footprint, and a ``crs`` grid-mapping variable. Pass it to :meth:`stilt.Footprint.aggregate` as a target, or to other tools that read CF grids. Projected grids need ``pyproj``. """ import xarray as xr is_longlat = self.is_longlat x_centers, y_centers = self.axes x_dim, y_dim = ("lon", "lat") if is_longlat else ("x", "y") ds = xr.Dataset(coords={x_dim: x_centers, y_dim: y_centers}) ds.attrs["Conventions"] = "CF-1.8" ds["crs"] = xr.DataArray(0, attrs=_cf_grid_mapping_attrs(self.projection)) ds[x_dim].attrs.update(cf_axis_attrs(x_dim)) ds[y_dim].attrs.update(cf_axis_attrs(y_dim)) return ds
__all__ = ["Bounds", "Grid", "cf_axis_attrs"]